Lesson 110: Semantic Segmentation & U-Net

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

What this lesson covers

The lesson, slide by slide

1. Semantic Segmentation & U-Net from Scratch

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.

2. By the end of this lesson you can

Objectives

  1. Define semantic segmentation and identify the output tensor shape (B, C, H, W)
  2. Trace the U-Net encoder-decoder path and explain why skip connections use concatenation, not addition
  3. Implement three upsampling methods — bilinear interpolation, ConvTranspose2d, pixel shuffle — and state the output-size formula for each
  4. Implement Dice loss and Focal loss from scratch and explain when each beats pixel-wise cross-entropy
  5. Build TinyUNet end-to-end and verify input/output shapes match (B, n_classes, H, W)

3. What survived from Convex Optimization?

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.

4. Semantic segmentation — the task

Section

Part 1 of 4

5. Every pixel gets a label

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 \]

taskinputoutputloss
image classification(B,C,H,W)(B, n_cls)CrossEntropyLoss
object detection(B,C,H,W)bounding boxesSmooth L1 + CE
semantic segmentation(B,C,H,W)(B, n_cls, H, W)pixel-wise CE / Dice

6. Fill in: input for Every pixel gets a label

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.

taskinputoutputloss
image classification(B,C,H,W)(B, n_cls)CrossEntropyLoss
object detection(B,C,H,W)bounding boxesSmooth L1 + CE
semantic segmentation(B,C,H,W)(B, n_cls, H, W)pixel-wise CE / Dice

7. Encoder-decoder: the core idea

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.

8. U-Net — skip connections & shapes

Section

Part 2 of 4

9. U-Net skip connections: concat, not add

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.

10. Guess the shape of the answer: U-Net tensor shapes — full trace

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.

11. U-Net tensor shapes — full trace

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)
stageshape (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.

12. What each one costs: U-Net tensor shapes — full trace

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.

stageshape (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

13. Something is wrong here: using addition instead of concatenation for U-Net skips

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.

14. Trap: using addition instead of concatenation for U-Net skips

Trap

The 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.

The fix

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.

15. Break it on purpose: using addition instead of concatenation for…

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.

16. Upsampling methods — three ways to grow H×W

Section

Part 3 of 4

17. Three upsampling methods compared

Concept

The decoder must recover spatial resolution lost during encoding. Three methods are standard; they differ in whether upsampling weights are fixed or learned.

methodlearnable paramshow it worksPyTorch
bilinear interpolationnone (fixed)weighted avg of 4 neighborsF.interpolate(..., mode='bilinear')
ConvTranspose2dyes (kernel)learned fractional-stride convnn.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).

18. Break it if you can: Three upsampling methods compared

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).

19. ConvTranspose2d — output size formula

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_Hformula (h-1)*2+2actual out_H
4(4-1)*2+2 = 88
8(8-1)*2+2 = 1616
16(16-1)*2+2 = 3232

20. All three upsamplers — verified shapes

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)
methodinput shapeoutput shapeC_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

21. Where does each piece belong: Lesson 110: Semantic Segmentation & U-Net

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.

Semantic segmentation — the task
Every pixel gets a label; Encoder-decoder: the core idea
U-Net — skip connections & shapes
U-Net skip connections: concat, not add; U-Net tensor shapes — full trace
Upsampling methods — three ways to grow H×W
Three upsampling methods compared; ConvTranspose2d — output size formula; All three upsamplers — verified shapes
s1
Semantic segmentation — the task is where Lesson 110: Semantic Segmentation & U-Net puts Every pixel gets a label, Encoder-decoder: the core idea. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
U-Net — skip connections & shapes is where Lesson 110: Semantic Segmentation & U-Net puts U-Net skip connections: concat, not add, U-Net tensor shapes — full trace. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
Upsampling methods — three ways to grow H×W is where Lesson 110: Semantic Segmentation & U-Net puts Three upsampling methods compared, ConvTranspose2d — output size formula, All three upsamplers — verified shapes. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

22. Something is wrong here: checkerboard artifacts from ConvTranspose2d

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.

23. Trap: checkerboard artifacts from ConvTranspose2d

Trap

The 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.

The fix

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.

24. Which of these survive contact with Lesson 110: Semantic Segmentation & U-Net?

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.

Holds up
The decoder must recover spatial resolution lost during encoding. Three methods are standard; they differ in whether upsampling weights are fixed or learned.; The Dice coefficient (also called F1 score for sets) measures overlap between predicted and ground-truth masks. Dice loss = 1 − Dice, averaged over classes.
Breaks
U-Net skip connections add the encoder feature map to the decoder: dec = upsample(prev) + enc_skip — same pattern as ResNet residuals (Lesson 70).; Use ConvTranspose2d for all upsampling steps and expect smooth outputs.
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 110: Semantic Segmentation & U-Net puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

25. Segmentation losses — CE, Dice, Focal

Section

Part 4 of 4

26. Why pixel-wise CE fails on imbalanced masks

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.

scenarioCE lossDice losswhich to trust
predict all background (1/2048 pixels are class-1)0.0092 (low!)0.5013Dice — CE is blind
random logits1.38740.6606both penalize
near-perfect prediction~0.0~0.0both 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.

27. Dice loss — overlap-based

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.

28. By analogy: Dice loss — overlap-based

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.

29. Dice loss from scratch

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))
inputdice_lossnote
random logits0.6606B=2,C=3,8x8, seed 7
near-perfect (*10)0.000092epsilon prevents exact 0

30. Fill in: dice_loss for Dice loss from scratch

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.

inputdice_lossnote
random logits0.6606B=2,C=3,8x8, seed 7
near-perfect (*10)0.000092epsilon prevents exact 0

31. Focal loss — hard-example mining

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
0standard CE — every pixel weighted equally
1mild down-weighting of easy pixels
2standard focal (used in RetinaNet); random logits give 0.9373 vs CE 1.3874
5very aggressive — focus almost entirely on the hardest examples

32. Focal loss from scratch

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}')
lossrandom logitsnear-perfectnote
CE (γ=0)1.3874~0.0all pixels equal weight
Focal (γ=2)0.9373~0.0down-weights easy pixels; hard ones dominate

33. Without one step: The semantic segmentation recipe

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:

  1. Output shape: (B, n_classes, H, W) — one logit map per class, same spatial size as input
  2. Encoder: repeated DoubleConv + MaxPool; channels grow, H/W shrink by 2 each stage
  3. Bottleneck: deepest DoubleConv — maximum context, minimum H/W
  4. Decoder: ConvTranspose2d(k=2,s=2) doubles H/W; then torch.cat([up, enc_skip], dim=1) + DoubleConv
  5. Skip connections: concat encoder feature map to matching decoder level (doubles channels); never add
  6. Output: Conv2d(base, n_classes, kernel_size=1) — 1×1 conv maps channels to class logits
  7. Loss: pixel-wise CE for balanced classes; Dice (overlap) or Focal (hard-example) for imbalanced masks; Dice + CE combined is common
  8. ConvTranspose2d size: H_out = (H_in - 1)*stride - 2*padding + kernel_size

34. The semantic segmentation recipe

Pattern

  1. Output shape: (B, n_classes, H, W) — one logit map per class, same spatial size as input
  2. Encoder: repeated DoubleConv + MaxPool; channels grow, H/W shrink by 2 each stage
  3. Bottleneck: deepest DoubleConv — maximum context, minimum H/W
  4. Decoder: ConvTranspose2d(k=2,s=2) doubles H/W; then torch.cat([up, enc_skip], dim=1) + DoubleConv
  5. Skip connections: concat encoder feature map to matching decoder level (doubles channels); never add
  6. Output: Conv2d(base, n_classes, kernel_size=1) — 1×1 conv maps channels to class logits
  7. Loss: pixel-wise CE for balanced classes; Dice (overlap) or Focal (hard-example) for imbalanced masks; Dice + CE combined is common
  8. ConvTranspose2d size: H_out = (H_in - 1)*stride - 2*padding + kernel_size

35. Where does it stop working: The semantic segmentation recipe

Edge 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:

  1. Output shape: (B, n_classes, H, W) — one logit map per class, same spatial size as input
  2. Encoder: repeated DoubleConv + MaxPool; channels grow, H/W shrink by 2 each stage
  3. Bottleneck: deepest DoubleConv — maximum context, minimum H/W
  4. Decoder: ConvTranspose2d(k=2,s=2) doubles H/W; then torch.cat([up, enc_skip], dim=1) + DoubleConv
  5. Skip connections: concat encoder feature map to matching decoder level (doubles channels); never add
  6. Output: Conv2d(base, n_classes, kernel_size=1) — 1×1 conv maps channels to class logits
  7. Loss: pixel-wise CE for balanced classes; Dice (overlap) or Focal (hard-example) for imbalanced masks; Dice + CE combined is common
  8. ConvTranspose2d size: H_out = (H_in - 1)*stride - 2*padding + kernel_size

36. Rule out three: Check yourself — skip connection type

Elimination

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.

  • A. torch.cat([upsample(dec_prev), enc_i], dim=1) — doubles the channel count before the decoder conv
  • B. upsample(dec_prev) + enc_i — element-wise addition, same as ResNet residuals
  • C. nn.Linear applied to the flattened encoder feature map
  • D. enc_i is used only to initialize the decoder weights, not at forward-pass time

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).

37. Check yourself — skip connection type

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:

  • A. torch.cat([upsample(dec_prev), enc_i], dim=1) — doubles the channel count before the decoder conv (correct)
  • B. upsample(dec_prev) + enc_i — element-wise addition, same as ResNet residuals
  • C. nn.Linear applied to the flattened encoder feature map
  • D. enc_i is used only to initialize the decoder weights, not at forward-pass time

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).

Why B tempts people
Element-wise addition (ResNet style) requires both tensors to have the same number of channels and forces their representations to align, discarding the independent identity of encoder vs decoder features. This is architecturally correct for ResNets but wrong for U-Net.
Why C tempts people
U-Net operates in the spatial domain — no flattening occurs. Spatial feature maps are concatenated, not collapsed to vectors.
Why D tempts people
The encoder skip feature maps are used at every forward pass at inference time, not just during initialization. They carry fine-grained spatial detail the bottleneck has lost.

38. Answer it before you see the options: Check yourself — ConvTranspose2d output…

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.

39. Check yourself — ConvTranspose2d output size

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?

  • A. (B, 16, 16, 16) (correct)
  • B. (B, 16, 8, 8)
  • C. (B, 32, 16, 16)
  • D. (B, 16, 4, 4)

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.

Why B tempts people
H=8 would require stride=1 (no upsampling). ConvTranspose2d with stride=2 is specifically designed to increase spatial dimensions by the stride factor.
Why C tempts people
C_out is set by the second constructor argument (16), not inherited from C_in (32). The spatial calculation H=16 is correct, but the channel count is wrong.
Why D tempts people
H=4 would require stride=0.5, which is not possible. ConvTranspose2d with stride=2 doubles (not halves) the spatial dimensions.

40. Rule out three: Check yourself — Dice vs CE for imbalanced masks

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.

  • A. CE is deceptively small (0.009) because 2047/2048 predictions are correct; Dice is large (0.50) because the foreground class has zero intersection
  • B. Dice is deceptively small because it averages over pixels; CE is large because it penalizes the missed foreground pixel heavily
  • C. Both losses are equally small; class imbalance only matters for accuracy metrics, not differentiable losses
  • D. Focal loss (γ=2) is deceptively small while both CE and Dice correctly penalize

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.

41. Check yourself — Dice vs CE for imbalanced masks

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?

  • A. CE is deceptively small (0.009) because 2047/2048 predictions are correct; Dice is large (0.50) because the foreground class has zero intersection (correct)
  • B. Dice is deceptively small because it averages over pixels; CE is large because it penalizes the missed foreground pixel heavily
  • C. Both losses are equally small; class imbalance only matters for accuracy metrics, not differentiable losses
  • D. Focal loss (γ=2) is deceptively small while both CE and Dice correctly penalize

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.

Why B tempts people
Dice averages over classes (not pixels), so the minority class has equal weight to the background class. CE is the one that effectively averages over pixels, making it blind to the rare foreground.
Why C tempts people
Class imbalance directly affects differentiable losses. The verified numbers (CE=0.009, Dice=0.50) show a 54× difference in signal strength for the same prediction.
Why D tempts people
Focal loss also suffers from imbalance when the model is confidently wrong on the background (high p_t for background → low focal weight). Dice's per-class overlap formulation is the most principled fix.

42. Your turn: TinyUNet from scratch

Section

Project

43. Project: TinyUNet on synthetic masks

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.

#milestonewhat to verify
1Implement TinyUNet, print input/output shapesoutput must be (B, 2, H, W) = input spatial size
2Implement dice_loss from scratchperfect logits → ~0; random logits → ~0.66
3Train 5 epochs with Dice + CE; print loss tracetotal 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.

44. Break it if you can: Project: TinyUNet on synthetic masks

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.

45. Milestone 1 — TinyUNet shape trace

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)
propertyvaluenote
total parameters117,410verified 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

46. Milestone 2 — Dice loss verified

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))
casedice_lossinterpretation
random logits (seed 7)0.6606above 0.5: more miss than hit
near-perfect (*10)0.000092epsilon prevents divide-by-zero

47. Milestone 3 — train with Dice + CE

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}')
epochCEDicetotal
00.79160.67921.4708
10.74960.67331.4229
20.72320.66881.3920
30.70170.66481.3665
40.68300.66111.3441

48. What each one costs: Milestone 3 — train with Dice + CE

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.

epochCEDicetotal
00.79160.67921.4708
10.74960.67331.4229
20.72320.66881.3920
30.70170.66481.3665
40.68300.66111.3441

49. The full program

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 linevalue
ep0 total1.4708
ep1 total1.4229
ep4 total1.3441
output shapetorch.Size([2, 2, 32, 32])

50. Fill in: value for The full program

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 linevalue
ep0 total1.4708
ep1 total1.4229
ep4 total1.3441
output shapetorch.Size([2, 2, 32, 32])

51. Show it off

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.

52. Connect it up: Lesson 110: Semantic Segmentation & U-Net

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.

53. What you can do now

Recap

conceptthe one thing to remember
segmentation output(B, n_cls, H, W) — same H/W as input
U-Net skipconcat on dim=1, not add — doubles channels
ConvTranspose2dH_out = (H_in-1)*s + k; stride=2, k=2 doubles H/W
Dice loss1 - mean_class_overlap; robust to imbalance
CE vs Dice imbalanceCE 0.009 vs Dice 0.50 predicting all-background

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 110 — Semantic Segmentation / U-Net — Barron · USAAIO Round 2 Preparation, 2026
  2. Ronneberger et al. 'U-Net: Convolutional Networks for Biomedical Image Segmentation' (MICCAI 2015) — arXiv:1505.04597
  3. Milletari et al. 'V-Net: Fully Convolutional Neural Networks for Volumetric Medical Image Segmentation' — Dice loss formulation — arXiv:1606.04797
  4. Lin et al. 'Focal Loss for Dense Object Detection' (RetinaNet, ICCV 2017) — arXiv:1708.02002
  5. TinyUNet from scratch, Dice/CE/Focal loss, ConvTranspose2d shapes, skip-connection concat vs add — all values verified with torch 2.7.1+cpu and numpy 2.2.6, June 2026 — Real execution, verified

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108