hadrilec's picture
description v0.1 vs v0
4175b07
Raw
History Blame
4.38 kB
import torch.nn as nn
import torch.nn.functional as F
class Net(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 16, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(16, 8, kernel_size=3, padding=1)
self.fc1 = nn.Linear(8 * 64 * 64, 32) # for 256x256 input
self.fc2 = nn.Linear(32, 4) # 4 classes
def forward(self, x):
out = F.max_pool2d(torch.tanh(self.conv1(x)), 2) # 256 -> 128
out = F.max_pool2d(torch.tanh(self.conv2(out)), 2) # 128 -> 64
out = out.view(-1, 8 * 64 * 64)
out = torch.tanh(self.fc1(out))
out = self.fc2(out)
return out
# Why moving from v0 to v0.1?
# 🔹 1. More convolution layers with increasing channels (16 → 32 → 64)
# Why more layers?
# Convolutions extract features in a hierarchy:
# Early layers → detect edges, corners, basic color patterns.
# Middle layers → detect textures, shapes (e.g., “forest canopy patterns”, “rooftop grids”).
# Deeper layers → detect high-level structures (urban blocks vs farmland).
# Your original network had just 2 layers (16, 8 channels) → too shallow to capture rich structures.
# Why increasing channels?
# Each channel is like a “detector” for a specific feature.
# With 16 channels, you only learn 16 low-level filters → very limited.
# Expanding 16 → 32 → 64 gives the network more capacity to encode complex patterns without exploding the number of parameters.
# Standard practice in CNNs: double channels as you go deeper because spatial resolution shrinks (e.g., 256 → 128 → 64 → 32 pixels).
# 📌 Effect: Model can distinguish “forest vs field” based not just on green color, but also on texture differences (dense tree canopy vs plowed farmland).
# 🔹 2. Replace tanh with ReLU
# tanh: outputs values between -1 and 1.
# Problem: saturates easily → gradients vanish when activations go to extremes.
# Slower learning, harder to train deeper nets.
# ReLU (Rectified Linear Unit): outputs 0 for negatives, identity for positives.
# Keeps gradients alive → avoids vanishing gradient problem.
# Sparse activations → only some neurons fire, which acts as implicit regularization.
# Faster and easier optimization, now the standard in almost all CNNs.
# 📌 Effect: Your network trains faster, learns stronger nonlinear features, and avoids being “stuck” in flat gradient regions.
# 🔹 3. Insert Batch Normalization (BatchNorm)
# BatchNorm normalizes activations (zero mean, unit variance) inside the network.
# Benefits:
# Stabilizes training → network is less sensitive to weight initialization.
# Allows higher learning rates → speeds up convergence.
# Acts as a regularizer, reducing overfitting.
# Makes the network more robust to distribution shifts (helpful since Google Earth tiles differ from IGN tiles).
# 📌 Effect: Your model generalizes better to “out-of-distribution” inputs, since it isn’t too sensitive to raw intensity/color shifts.
# 🔹 4. Add Dropout before fully connected head
# Dropout randomly “drops” (sets to zero) some neurons during training.
# Forces the network not to rely on a single pathway → encourages redundancy in feature representation.
# Prevents co-adaptation of neurons (where certain features only work if others are present).
# Especially important in the fully connected head, which has many parameters relative to convolutional layers.
# 📌 Effect: Reduces overfitting to IGN imagery and makes the network more robust when faced with slightly different textures/colors from Google Earth.
# 🔹 5. Keep output layer at 4 classes
# Since your problem is 4-way classification (sea, forest, urban, field), final layer must output 4 logits.
# Logits go into nn.CrossEntropyLoss, which applies softmax internally.
# 📌 Effect: Keeps task-specific output while allowing the backbone (convs + FCs) to be stronger.
# ✅ Summary of why this matters for you
# Your old net was too shallow, used tanh (harder training), no normalization (unstable activations), and no dropout (overfits to IGN style). With these fixes, you’ll:
# Learn richer features (textures, structures).
# Train faster & more stable (ReLU + BatchNorm).
# Generalize better (Dropout + normalization).
# Stay lightweight and trainable from scratch.