USAAIO Lesson 110, from Phase 3, on pixel-wise classification through semantic segmentation. It covers the U-Net encoder-decoder with its skip connections, which concatenate rather than add, and the upsampling methods available: bilinear interpolation, ConvTranspose2d, and pixel shuffle. It then covers Dice loss and Focal loss for class-imbalanced masks, and trains a TinyUNet from scratch. All the shapes, loss values, and training traces were verified with torch 2.7.1+cpu. The lesson runs to 32 slides.
Subject: Machine Learning · 53 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 110 · Phase 3
Assign a class label to every pixel. Build the encoder-decoder with skip connections, three upsampling methods, and two imbalance-aware losses — all in PyTorch, all verified.
Objectives
(B, C, H, W)ConvTranspose2d, pixel shuffle — and state the output-size formula for eachTinyUNet end-to-end and verify input/output shapes match (B, n_classes, H, W)Warm-up
Discussion prompt
Before we open Lesson 110: Semantic Segmentation & U-Net: without looking back, what was the main idea of Convex Optimization, and what could you do by the end of it that you could not do before?
Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.
Answer:
convex functions, the first- and second-order conditions, why convex problems have only global minima, gradient-descent convergence rates, and learning-rate selection via the smoothness bound. Build a GD convergence experiment from scratch.
Section
Part 1 of 4
Concept
Semantic segmentation is dense classification: instead of one label per image (Lessons 55-92), every pixel is assigned a class. The model output is a (B, C, H, W) tensor of per-class logits — one spatial map per class.
\[ \hat{y}_{b,c,i,j} = \text{logit for class } c \text{ at pixel } (i,j) \text{ in image } b \]
| task | input | output | loss |
|---|---|---|---|
| image classification | (B,C,H,W) | (B, n_cls) | CrossEntropyLoss |
| object detection | (B,C,H,W) | bounding boxes | Smooth L1 + CE |
| semantic segmentation | (B,C,H,W) | (B, n_cls, H, W) | pixel-wise CE / Dice |
Comparison
Comparison matrix
From Every pixel gets a label: refill the input column from what you know. The rest of the table is as it appeared.
| task | input | output | loss |
|---|---|---|---|
| image classification | (B,C,H,W) | (B, n_cls) | CrossEntropyLoss |
| object detection | (B,C,H,W) | bounding boxes | Smooth L1 + CE |
| semantic segmentation | (B,C,H,W) | (B, n_cls, H, W) | pixel-wise CE / Dice |
Concept
A plain CNN loses spatial resolution via pooling — useful for classification, fatal for segmentation. The encoder-decoder architecture restores it: downsample for context, then upsample to full resolution.
The problem: the decoder only sees a compressed bottleneck, so fine-grained edges and textures are lost. U-Net's solution is skip connections from each encoder level directly to the matching decoder level.
Section
Part 2 of 4
Concept
U-Net (Ronneberger 2015) connects each encoder stage to the same-resolution decoder stage with a concatenation skip connection — the encoder feature map is appended along the channel dimension.
\[ \text{dec}_{in} = \text{concat}(\text{upsample}(\text{dec}_{prev}),\; \text{enc}_{same}) \quad \text{(channel dim)} \]
This doubles the channel count going into the decoder conv block, so the network can independently learn what from the encoder and what from the decoder are useful — the ResNet addition (Lesson 70) would force them to the same representation.
Estimation
Predict first
Trace a U-Net with base=16 channels on a (B=2, 1, 32, 32) input. Predict each shape before reading on.
Commit before you compute: what does U-Net tensor shapes — full trace come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Concat(up_b, e2) along dim=1: (2,32,16,16)+(2,32,16,16) → (2,64,16,16)
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Skip concatenation doubles channels at this level — the decoder conv must accept 64 channels (32 from upsample + 32 from encoder skip), then reduce back to 32.
Worked example
Trace a U-Net with base=16 channels on a (B=2, 1, 32, 32) input. Predict each shape before reading on.
import torch, torch.nn as nn
class DoubleConv(nn.Module):
def __init__(self, in_c, out_c):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_c, out_c, 3, padding=1),
nn.BatchNorm2d(out_c), nn.ReLU(inplace=True),
nn.Conv2d(out_c, out_c, 3, padding=1),
nn.BatchNorm2d(out_c), nn.ReLU(inplace=True),
)
def forward(self, x): return self.conv(x)
torch.manual_seed(42)
x = torch.randn(2, 1, 32, 32)
enc1 = DoubleConv(1, 16); e1 = enc1(x) # no pool
enc2 = DoubleConv(16, 32); e2 = enc2(nn.MaxPool2d(2)(e1))
bottle = DoubleConv(32, 64); b = bottle(nn.MaxPool2d(2)(e2))
print('e1:', e1.shape, ' e2:', e2.shape, ' bottle:', b.shape)| stage | shape (B,C,H,W) | note |
|---|---|---|
| input | (2, 1, 32, 32) | B=2, C=1, H=W=32 |
| enc1 (2xConv) | (2, 16, 32, 32) | channels 1→16, no pool |
| pool → enc2 | (2, 32, 16, 16) | H/2; channels 16→32 |
| pool → bottle | (2, 64, 8, 8) | H/4; channels 32→64 |
ConvTranspose2d(64→32, k=2, s=2) upsample: (2,64,8,8) → (2,32,16,16)
Why: Output size = (in-1)stride + kernel = (8-1)2 + 2 = 16. The transposed conv both upsamples and halves channels.
Concat(up_b, e2) along dim=1: (2,32,16,16)+(2,32,16,16) → (2,64,16,16)
Why: Skip concatenation doubles channels at this level — the decoder conv must accept 64 channels (32 from upsample + 32 from encoder skip), then reduce back to 32.
Trade off
Comparison matrix
From U-Net tensor shapes — full trace: every row here is a choice with a cost. Fill the shape (B,C,H,W) column, then say which row you would actually pick and what you give up for it.
| stage | shape (B,C,H,W) | note |
|---|---|---|
| input | (2, 1, 32, 32) | B=2, C=1, H=W=32 |
| enc1 (2xConv) | (2, 16, 32, 32) | channels 1→16, no pool |
| pool → enc2 | (2, 32, 16, 16) | H/2; channels 16→32 |
| pool → bottle | (2, 64, 8, 8) | H/4; channels 32→64 |
Anomaly
Predict first
A student writes this, and it looks reasonable:
U-Net skip connections add the encoder feature map to the decoder: dec = upsample(prev) + enc_skip — same pattern as ResNet residuals (Lesson 70).
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Addition requires identical channel counts and forces the representations to align element-wise.
U-Net skip connections concatenate along the channel dimension: dec = torch.cat([upsample(prev), enc_skip], dim=1).
Why: Addition requires identical channel counts and forces the representations to align element-wise. This discards the identity of each source — the model can no longer distinguish encoder vs decoder features independently.
Trap
U-Net skip connections add the encoder feature map to the decoder: dec = upsample(prev) + enc_skip — same pattern as ResNet residuals (Lesson 70).
dec = upsample(prev) + enc_skip
Why: Addition requires identical channel counts and forces the representations to align element-wise. This discards the identity of each source — the model can no longer distinguish encoder vs decoder features independently.
U-Net skip connections concatenate along the channel dimension: dec = torch.cat([upsample(prev), enc_skip], dim=1).
dec = torch.cat([upsample(prev), enc_skip], dim=1)
Why: Concatenation preserves both representations independently — the following DoubleConv learns which features from each source are useful. This is the defining architectural choice that separates U-Net from plain encoder-decoders.
Break the constraint
Discussion prompt
The rule this trap just fixed:
U-Net skip connections concatenate along the channel dimension: dec = torch.cat([upsample(prev), enc_skip], dim=1).
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
Addition requires identical channel counts and forces the representations to align element-wise. This discards the identity of each source — the model can no longer distinguish encoder vs decoder features independently.
Section
Part 3 of 4
Concept
The decoder must recover spatial resolution lost during encoding. Three methods are standard; they differ in whether upsampling weights are fixed or learned.
| method | learnable params | how it works | PyTorch |
|---|---|---|---|
| bilinear interpolation | none (fixed) | weighted avg of 4 neighbors | F.interpolate(..., mode='bilinear') |
| ConvTranspose2d | yes (kernel) | learned fractional-stride conv | nn.ConvTranspose2d(C_in, C_out, k=2, s=2) |
| pixel shuffle (sub-pixel conv) | yes (in conv before) | rearrange (B,C·r²,H,W) → (B,C,rH,rW) | nn.PixelShuffle(r) |
U-Net canonically uses ConvTranspose2d. Bilinear + Conv is also popular (avoids checkerboard artifacts). Pixel shuffle is common in super-resolution (Lesson 98 territory).
Counterexample
Discussion prompt
The decoder must recover spatial resolution lost during encoding. Three methods are standard; they differ in whether upsampling weights are fixed or learned.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
U-Net canonically uses ConvTranspose2d. Bilinear + Conv is also popular (avoids checkerboard artifacts). Pixel shuffle is common in super-resolution (Lesson 98 territory).
Worked example
Predict the output size of ConvTranspose2d(4, 2, kernel_size=2, stride=2) applied to inputs of H=4, 8, 16. Use the formula before running.
\[ H_{out} = (H_{in} - 1) \cdot \text{stride} - 2 \cdot \text{padding} + \text{kernel\_size} \]
import torch, torch.nn as nn
ct = nn.ConvTranspose2d(4, 2, kernel_size=2, stride=2)
for h in [4, 8, 16]:
x = torch.randn(1, 4, h, h)
with torch.no_grad():
out = ct(x)
formula = (h - 1) * 2 - 0 + 2 # padding=0
print(f'in={h:>2} -> out={out.shape[-1]:>2} (formula {formula})')| in_H | formula (h-1)*2+2 | actual out_H |
|---|---|---|
| 4 | (4-1)*2+2 = 8 | 8 |
| 8 | (8-1)*2+2 = 16 | 16 |
| 16 | (16-1)*2+2 = 32 | 32 |
Worked example
Apply all three methods to the same (2, 8, 4, 4) tensor and confirm shapes. Pixel shuffle requires a preceding conv to expand channels to C * r².
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(42)
x = torch.randn(2, 8, 4, 4)
# Bilinear (no params)
up_bil = F.interpolate(x, scale_factor=2, mode='bilinear', align_corners=False)
print('bilinear: ', up_bil.shape)
# ConvTranspose2d (learned)
t_conv = nn.ConvTranspose2d(8, 4, kernel_size=2, stride=2)
up_ct = t_conv(x)
print('ConvTranspose:', up_ct.shape)
# PixelShuffle r=2: need (B, C*r^2, H, W) = (B, 8*4=32, 4, 4)
x_ps = torch.randn(2, 32, 4, 4)
up_ps = nn.PixelShuffle(2)(x_ps)
print('PixelShuffle: ', up_ps.shape)| method | input shape | output shape | C_out |
|---|---|---|---|
| bilinear | (2, 8, 4, 4) | (2, 8, 8, 8) | 8 (unchanged) |
| ConvTranspose2d | (2, 8, 4, 4) | (2, 4, 8, 8) | 4 (set by C_out) |
| PixelShuffle r=2 | (2, 32, 4, 4) | (2, 8, 8, 8) | C/r² = 32/4 = 8 |
Sorting
Sort into buckets
These are the pieces of Lesson 110: Semantic Segmentation & U-Net, out of order. Put each one back under the part of the lesson it belongs to.
Anomaly
Predict first
A student writes this, and it looks reasonable:
Use ConvTranspose2d for all upsampling steps and expect smooth outputs.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Learned transposed convolutions with stride=2 can produce checkerboard artifacts — high-frequency grid patterns in the output — because overlapping kernel footprints receive unequal gradient updates during training.
Use bilinear upsample followed by a regular Conv2d — this avoids checkerboard patterns entirely.
Why: Learned transposed convolutions with stride=2 can produce checkerboard artifacts — high-frequency grid patterns in the output — because overlapping kernel footprints receive unequal gradient updates during training.
Trap
Use ConvTranspose2d for all upsampling steps and expect smooth outputs.
Use ConvTranspose2d with large strides for all decoder upsampling
Why: Learned transposed convolutions with stride=2 can produce checkerboard artifacts — high-frequency grid patterns in the output — because overlapping kernel footprints receive unequal gradient updates during training.
Use bilinear upsample followed by a regular Conv2d — this avoids checkerboard patterns entirely.
F.interpolate(..., mode='bilinear') then nn.Conv2d(C, C, 3, padding=1)
Why: Bilinear interpolation is smooth by definition. The subsequent Conv2d learns local refinements without the uneven footprint problem. Many modern U-Net variants (e.g., segmentation-models-pytorch) use this pattern by default.
Two truths and a lie
Sort into buckets
Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.
dec = upsample(prev) + enc_skip — same pattern as ResNet residuals (Lesson 70).; Use ConvTranspose2d for all upsampling steps and expect smooth outputs.Section
Part 4 of 4
Concept
Pixel-wise CrossEntropyLoss treats every pixel equally. In medical/satellite images, the foreground class may be 1% of pixels — predicting background everywhere gives 99% accuracy while completely missing the object.
| scenario | CE loss | Dice loss | which to trust |
|---|---|---|---|
| predict all background (1/2048 pixels are class-1) | 0.0092 (low!) | 0.5013 | Dice — CE is blind |
| random logits | 1.3874 | 0.6606 | both penalize |
| near-perfect prediction | ~0.0 | ~0.0 | both agree |
Numbers above are from TinyUNet verification on a (2, 2, 32, 32) mask where only 1/2048 pixels are class-1 — CE barely moves while Dice correctly penalizes the missed minority.
Concept
The Dice coefficient (also called F1 score for sets) measures overlap between predicted and ground-truth masks. Dice loss = 1 − Dice, averaged over classes.
\[ \text{Dice}(P, G) = \frac{2 \sum_{i} p_i g_i}{\sum_i p_i + \sum_i g_i + \varepsilon} \qquad \text{Dice loss} = 1 - \frac{1}{C}\sum_{c=1}^C \text{Dice}_c \]
P is the softmax probability map, G is the one-hot ground-truth map, ε prevents division by zero. Range: 0 (perfect) to 1 (no overlap). Naturally handles class imbalance because small classes have the same weight as large ones.
Analogy
Discussion prompt
Explain Dice loss — overlap-based by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
The Dice coefficient (also called F1 score for sets) measures overlap between predicted and ground-truth masks. Dice loss = 1 − Dice, averaged over classes.
Worked example
Implement Dice loss and verify on random and near-perfect logits. Predict whether the perfect case reaches exactly 0.
import torch, torch.nn as nn, torch.nn.functional as F
def dice_loss(pred, target, eps=1e-6):
"""pred: (B, C, H, W) logits; target: (B, H, W) int labels"""
B, C, H, W = pred.shape
pred_soft = F.softmax(pred, dim=1) # (B, C, H, W)
target_oh = F.one_hot(target, num_classes=C) # (B, H, W, C)
target_oh = target_oh.permute(0, 3, 1, 2).float()
intersection = (pred_soft * target_oh).sum(dim=[0, 2, 3])
denom = (pred_soft + target_oh).sum(dim=[0, 2, 3])
dice_per_cls = (2 * intersection + eps) / (denom + eps)
return 1.0 - dice_per_cls.mean()
torch.manual_seed(7)
B, C, H, W = 2, 3, 8, 8
pred = torch.randn(B, C, H, W)
target = torch.randint(0, C, (B, H, W))
print('random:', round(dice_loss(pred, target).item(), 4))
perf = F.one_hot(target, C).permute(0,3,1,2).float() * 10
print('perfect:', round(dice_loss(perf, target).item(), 6))| input | dice_loss | note |
|---|---|---|
| random logits | 0.6606 | B=2,C=3,8x8, seed 7 |
| near-perfect (*10) | 0.000092 | epsilon prevents exact 0 |
Comparison
Comparison matrix
From Dice loss from scratch: refill the dice_loss column from what you know. The rest of the table is as it appeared.
| input | dice_loss | note |
|---|---|---|
| random logits | 0.6606 | B=2,C=3,8x8, seed 7 |
| near-perfect (*10) | 0.000092 | epsilon prevents exact 0 |
Concept
Focal loss (Lin 2017, RetinaNet) down-weights easy examples by multiplying the CE by (1 − p_t)^γ where p_t is the predicted probability for the correct class. High confidence → small weight; hard examples → full weight.
\[ \text{FL}(p_t) = -(1-p_t)^\gamma \log(p_t) \qquad \gamma=0 \Rightarrow \text{standard CE} \]
| γ | effect |
|---|---|
| 0 | standard CE — every pixel weighted equally |
| 1 | mild down-weighting of easy pixels |
| 2 | standard focal (used in RetinaNet); random logits give 0.9373 vs CE 1.3874 |
| 5 | very aggressive — focus almost entirely on the hardest examples |
Worked example
Implement focal_loss(gamma=2) and verify it is strictly larger than CE on random logits (easy + hard mixed), and falls near 0 on near-perfect logits. Explain the direction.
import torch, torch.nn.functional as F
def focal_loss(pred, target, gamma=2.0):
"""pred: (B,C,H,W) logits; target: (B,H,W) int labels"""
B, C, H, W = pred.shape
log_p = F.log_softmax(pred, dim=1) # (B,C,H,W)
p = log_p.exp()
log_pt = log_p.gather(1, target.unsqueeze(1)).squeeze(1) # (B,H,W)
pt = p.gather(1, target.unsqueeze(1)).squeeze(1)
return -(((1 - pt) ** gamma) * log_pt).mean()
torch.manual_seed(7)
B, C, H, W = 2, 3, 8, 8
pred_rnd = torch.randn(B, C, H, W)
target = torch.randint(0, C, (B, H, W))
ce_rnd = F.cross_entropy(pred_rnd, target)
fl_rnd = focal_loss(pred_rnd, target, gamma=2)
print(f'CE random: {ce_rnd.item():.4f}')
print(f'Focal(g=2) rnd:{fl_rnd.item():.4f}')
perf = F.one_hot(target,C).permute(0,3,1,2).float()*20
print(f'Focal(g=2) perfect: {focal_loss(perf, target).item():.6f}')| loss | random logits | near-perfect | note |
|---|---|---|---|
| CE (γ=0) | 1.3874 | ~0.0 | all pixels equal weight |
| Focal (γ=2) | 0.9373 | ~0.0 | down-weights easy pixels; hard ones dominate |
Constraint
Discussion prompt
Run The semantic segmentation recipe with this step confiscated:
Skip connections: concat encoder feature map to matching decoder level (doubles channels); never add
Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.
Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.
Answer:
(B, n_classes, H, W) — one logit map per class, same spatial size as inputConvTranspose2d(k=2,s=2) doubles H/W; then torch.cat([up, enc_skip], dim=1) + DoubleConvConv2d(base, n_classes, kernel_size=1) — 1×1 conv maps channels to class logitsH_out = (H_in - 1)*stride - 2*padding + kernel_sizePattern
(B, n_classes, H, W) — one logit map per class, same spatial size as inputConvTranspose2d(k=2,s=2) doubles H/W; then torch.cat([up, enc_skip], dim=1) + DoubleConvConv2d(base, n_classes, kernel_size=1) — 1×1 conv maps channels to class logitsH_out = (H_in - 1)*stride - 2*padding + kernel_sizeEdge cases
Discussion prompt
The semantic segmentation recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
(B, n_classes, H, W) — one logit map per class, same spatial size as inputConvTranspose2d(k=2,s=2) doubles H/W; then torch.cat([up, enc_skip], dim=1) + DoubleConvConv2d(base, n_classes, kernel_size=1) — 1×1 conv maps channels to class logitsH_out = (H_in - 1)*stride - 2*padding + kernel_sizeElimination
Eliminate the wrong options
In U-Net, the skip connection from encoder level i to decoder level i uses:
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: U-Net concatenates the encoder feature map to the upsampled decoder tensor along the channel dimension (dim=1). This doubles channels — e.g., (B,32,H,W) cat (B,32,H,W) = (B,64,H,W) — and lets the following DoubleConv learn independently from both sources. Verified in the shape trace: concat shapes are (2,32,16,16)+(2,32,16,16)=(2,64,16,16).
Check
Identify the structural difference before clicking.
Check your understanding
In U-Net, the skip connection from encoder level i to decoder level i uses:
Answer: A
Why: U-Net concatenates the encoder feature map to the upsampled decoder tensor along the channel dimension (dim=1). This doubles channels — e.g., (B,32,H,W) cat (B,32,H,W) = (B,64,H,W) — and lets the following DoubleConv learn independently from both sources. Verified in the shape trace: concat shapes are (2,32,16,16)+(2,32,16,16)=(2,64,16,16).
Prediction
Predict first
nn.ConvTranspose2d(32, 16, kernel_size=2, stride=2, padding=0) is applied to a (B, 32, 8, 8) tensor. What is the output shape?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: (B, 16, 16, 16)
Why: H_out = (8-1)2 - 20 + 2 = 14 + 2 = 16. C_out = 16 (set by the second argument). Full shape: (B, 16, 16, 16). This is the standard U-Net decoder step: stride=2 and kernel_size=2 together double the spatial dimensions.
Check
Apply the formula before running.
Check your understanding
nn.ConvTranspose2d(32, 16, kernel_size=2, stride=2, padding=0) is applied to a (B, 32, 8, 8) tensor. What is the output shape?
Answer: A
Why: H_out = (8-1)2 - 20 + 2 = 14 + 2 = 16. C_out = 16 (set by the second argument). Full shape: (B, 16, 16, 16). This is the standard U-Net decoder step: stride=2 and kernel_size=2 together double the spatial dimensions.
Elimination
Eliminate the wrong options
A model predicts background for every pixel. The ground-truth mask has 1 foreground pixel out of 2048. Which loss value is deceptively small, and why?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: Verified numerically: predicting all class-0 with only 1/2048 pixels being class-1 gives CE=0.0092 (dominated by the 2047 correct background pixels) and Dice=0.5013 (foreground class has intersection=0, so Dice_fg ≈ 0, Dice_bg ≈ 1, mean = 0.5, loss = 1-0.5 = 0.5). Dice is imbalance-robust because it measures overlap per class, not per pixel.
Check
Reason from the loss formulas.
Check your understanding
A model predicts background for every pixel. The ground-truth mask has 1 foreground pixel out of 2048. Which loss value is deceptively small, and why?
Answer: A
Why: Verified numerically: predicting all class-0 with only 1/2048 pixels being class-1 gives CE=0.0092 (dominated by the 2047 correct background pixels) and Dice=0.5013 (foreground class has intersection=0, so Dice_fg ≈ 0, Dice_bg ≈ 1, mean = 0.5, loss = 1-0.5 = 0.5). Dice is imbalance-robust because it measures overlap per class, not per pixel.
Section
Project
Concept
Build TinyUNet (1 input channel, 2 classes, base=16) and train it with combined Dice + CE loss on a synthetic segmentation dataset. Three milestones: shape trace → implement Dice loss → train the model.
| # | milestone | what to verify |
|---|---|---|
| 1 | Implement TinyUNet, print input/output shapes | output must be (B, 2, H, W) = input spatial size |
| 2 | Implement dice_loss from scratch | perfect logits → ~0; random logits → ~0.66 |
| 3 | Train 5 epochs with Dice + CE; print loss trace | total loss should decrease each epoch |
Build rules: type every line; print shapes after every concat; check logits.shape[-2:] == x.shape[-2:] before computing loss; verify Dice on a trivial case before plugging into the training loop.
Counterexample
Discussion prompt
Build rules: type every line; print shapes after every concat; check logits.shape[-2:] == x.shape[-2:] before computing loss; verify Dice on a trivial case before plugging into the training loop.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Worked example
Your turn: implement TinyUNet and verify the output shape. Before coding, predict: after the two ConvTranspose2d upsample steps, does the spatial size return to 32×32?
Hint: use two encoder levels (base=16, base2=32) and one bottleneck (base4=64). Each up step: ConvTranspose2d(in, out, k=2, s=2) then torch.cat([up, enc_skip], dim=1) then DoubleConv.
import torch, torch.nn as nn
class DoubleConv(nn.Module):
def __init__(self, in_c, out_c):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_c, out_c, 3, padding=1), nn.BatchNorm2d(out_c), nn.ReLU(inplace=True),
nn.Conv2d(out_c, out_c, 3, padding=1), nn.BatchNorm2d(out_c), nn.ReLU(inplace=True),
)
def forward(self, x): return self.conv(x)
class TinyUNet(nn.Module):
def __init__(self, in_c=1, n_classes=2, base=16):
super().__init__()
self.enc1 = DoubleConv(in_c, base)
self.pool1 = nn.MaxPool2d(2)
self.enc2 = DoubleConv(base, base*2)
self.pool2 = nn.MaxPool2d(2)
self.bottleneck = DoubleConv(base*2, base*4)
self.up2 = nn.ConvTranspose2d(base*4, base*2, kernel_size=2, stride=2)
self.dec2 = DoubleConv(base*4, base*2) # cat: base*2 + base*2
self.up1 = nn.ConvTranspose2d(base*2, base, kernel_size=2, stride=2)
self.dec1 = DoubleConv(base*2, base) # cat: base + base
self.out_conv = nn.Conv2d(base, n_classes, kernel_size=1)
def forward(self, x):
e1 = self.enc1(x)
e2 = self.enc2(self.pool1(e1))
b = self.bottleneck(self.pool2(e2))
d2 = self.dec2(torch.cat([self.up2(b), e2], dim=1))
d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))
return self.out_conv(d1)
torch.manual_seed(42)
model = TinyUNet()
x = torch.randn(2, 1, 32, 32)
with torch.no_grad(): out = model(x)
print('params:', sum(p.numel() for p in model.parameters()))
print('output:', out.shape)| property | value | note |
|---|---|---|
| total parameters | 117,410 | verified with torch 2.7.1+cpu |
| input shape | (2, 1, 32, 32) | B=2, C=1 |
| output shape | (2, 2, 32, 32) | B=2, n_classes=2, same H/W |
Worked example
Your turn: implement dice_loss. Before coding, predict: does Dice = 1 when intersection = 0? What happens when pred_soft is uniform (all classes equally likely)?
Hint: F.softmax(pred, dim=1) gives probabilities; F.one_hot(target, C).permute(0,3,1,2).float() gives the ground-truth mask. Sum intersection and union over [0, 2, 3] (batch + spatial), then take (2*intersection + eps) / (union + eps) per class.
import torch, torch.nn.functional as F
def dice_loss(pred, target, eps=1e-6):
B, C, H, W = pred.shape
pred_soft = F.softmax(pred, dim=1)
target_oh = F.one_hot(target, C).permute(0, 3, 1, 2).float()
inter = (pred_soft * target_oh).sum(dim=[0, 2, 3])
denom = (pred_soft + target_oh).sum(dim=[0, 2, 3])
return 1.0 - ((2 * inter + eps) / (denom + eps)).mean()
torch.manual_seed(7)
pred = torch.randn(2, 3, 8, 8)
tgt = torch.randint(0, 3, (2, 8, 8))
perf = F.one_hot(tgt, 3).permute(0,3,1,2).float() * 10
print('random:', round(dice_loss(pred, tgt).item(), 4))
print('perfect:', round(dice_loss(perf, tgt).item(), 6))| case | dice_loss | interpretation |
|---|---|---|
| random logits (seed 7) | 0.6606 | above 0.5: more miss than hit |
| near-perfect (*10) | 0.000092 | epsilon prevents divide-by-zero |
Worked example
Your turn: train TinyUNet for 5 epochs with total_loss = CE + Dice. Before running, predict: should total loss start above 1.0 and decrease each epoch?
Hint: create a fixed synthetic target torch.zeros(B, H, W, dtype=torch.long) (all background). CE and Dice together give complementary gradients — CE from per-pixel signal, Dice from overlap.
import torch, torch.nn as nn
torch.manual_seed(42)
model = TinyUNet() # defined above
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
ce_fn = nn.CrossEntropyLoss()
x = torch.randn(2, 1, 32, 32)
target = torch.zeros(2, 32, 32, dtype=torch.long)
for ep in range(5):
opt.zero_grad()
out = model(x)
loss = ce_fn(out, target) + dice_loss(out, target)
loss.backward(); opt.step()
print(f'epoch {ep}: total={loss.item():.4f}')| epoch | CE | Dice | total |
|---|---|---|---|
| 0 | 0.7916 | 0.6792 | 1.4708 |
| 1 | 0.7496 | 0.6733 | 1.4229 |
| 2 | 0.7232 | 0.6688 | 1.3920 |
| 3 | 0.7017 | 0.6648 | 1.3665 |
| 4 | 0.6830 | 0.6611 | 1.3441 |
Trade off
Comparison matrix
From Milestone 3 — train with Dice + CE: every row here is a choice with a cost. Fill the Dice column, then say which row you would actually pick and what you give up for it.
| epoch | CE | Dice | total |
|---|---|---|---|
| 0 | 0.7916 | 0.6792 | 1.4708 |
| 1 | 0.7496 | 0.6733 | 1.4229 |
| 2 | 0.7232 | 0.6688 | 1.3920 |
| 3 | 0.7017 | 0.6648 | 1.3665 |
| 4 | 0.6830 | 0.6611 | 1.3441 |
Concept
import torch, torch.nn as nn, torch.nn.functional as F
class DoubleConv(nn.Module):
def __init__(self, i, o):
super().__init__()
self.c = nn.Sequential(
nn.Conv2d(i,o,3,padding=1), nn.BatchNorm2d(o), nn.ReLU(True),
nn.Conv2d(o,o,3,padding=1), nn.BatchNorm2d(o), nn.ReLU(True))
def forward(self, x): return self.c(x)
class TinyUNet(nn.Module):
def __init__(self, in_c=1, n_cls=2, b=16):
super().__init__()
self.e1=DoubleConv(in_c,b); self.p1=nn.MaxPool2d(2)
self.e2=DoubleConv(b,b*2); self.p2=nn.MaxPool2d(2)
self.bn=DoubleConv(b*2,b*4)
self.u2=nn.ConvTranspose2d(b*4,b*2,2,2); self.d2=DoubleConv(b*4,b*2)
self.u1=nn.ConvTranspose2d(b*2,b,2,2); self.d1=DoubleConv(b*2,b)
self.out=nn.Conv2d(b,n_cls,1)
def forward(self,x):
e1=self.e1(x); e2=self.e2(self.p1(e1)); bn=self.bn(self.p2(e2))
d2=self.d2(torch.cat([self.u2(bn),e2],1))
return self.out(self.d1(torch.cat([self.u1(d2),e1],1)))
def dice_loss(pred, target, eps=1e-6):
B,C,H,W=pred.shape
ps=F.softmax(pred,1); oh=F.one_hot(target,C).permute(0,3,1,2).float()
i=(ps*oh).sum([0,2,3]); d=(ps+oh).sum([0,2,3])
return 1-((2*i+eps)/(d+eps)).mean()
torch.manual_seed(42)
m=TinyUNet(); opt=torch.optim.Adam(m.parameters(),lr=1e-3); ce=nn.CrossEntropyLoss()
x=torch.randn(2,1,32,32); t=torch.zeros(2,32,32,dtype=torch.long)
for ep in range(5):
opt.zero_grad(); o=m(x)
loss=ce(o,t)+dice_loss(o,t); loss.backward(); opt.step()
print(f'ep{ep} total={loss.item():.4f}')
print('output shape:', m(x).shape) # torch.Size([2, 2, 32, 32])| output line | value |
|---|---|
| ep0 total | 1.4708 |
| ep1 total | 1.4229 |
| ep4 total | 1.3441 |
| output shape | torch.Size([2, 2, 32, 32]) |
Comparison
Comparison matrix
From The full program: refill the value column from what you know. The rest of the table is as it appeared.
| output line | value |
|---|---|
| ep0 total | 1.4708 |
| ep1 total | 1.4229 |
| ep4 total | 1.3441 |
| output shape | torch.Size([2, 2, 32, 32]) |
Concept
Out loud, slides closed: (1) trace the full TinyUNet forward pass from (2,1,32,32) input to (2,2,32,32) output, naming every shape at each encoder/bottleneck/decoder step; (2) explain why U-Net concatenates skip connections instead of adding them; (3) explain why Dice loss gives 0.50 when a model predicts all-background on a 1/2048 foreground mask, while CE gives only 0.009.
Stretch (homework): replace ConvTranspose2d with bilinear + Conv2d and compare outputs for checkerboard artifacts. Add Focal loss (γ=2) to the combined loss and compare convergence on an imbalanced mask (1% foreground). Try training on sklearn.datasets.load_digits reshaped as segmentation targets.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Semantic segmentation — the task · U-Net — skip connections & shapes · Upsampling methods — three ways to grow H×W · Segmentation losses — CE, Dice, Focal · Your turn: TinyUNet from scratch. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
(B, n_classes, H, W) and distinguish it from classificationConvTranspose2d output formula: H_out = (H_in - 1)*stride + kernel_size| concept | the one thing to remember |
|---|---|
| segmentation output | (B, n_cls, H, W) — same H/W as input |
| U-Net skip | concat on dim=1, not add — doubles channels |
| ConvTranspose2d | H_out = (H_in-1)*s + k; stride=2, k=2 doubles H/W |
| Dice loss | 1 - mean_class_overlap; robust to imbalance |
| CE vs Dice imbalance | CE 0.009 vs Dice 0.50 predicting all-background |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.