Lesson 58 of 60 · python
Convolutional Neural Networks (CNN) for Image Classification
Duration: 30 minutes
CNNs – Image Classification
Convolutional Neural Networks excel at extracting spatial hierarchies from images.
Key layers
- Conv2D: learns filters (kernels).
- MaxPooling2D: spatial down‑sampling.
- Flatten: convert 2‑D feature maps to 1‑D.
- Dropout: regularization.
Building a CNN on MNIST
import tensorflow as tf
from tensorflow.keras import layers, models
# Load data
mnist = tf.keras.datasets.mnist
(X_train, y_train), (X_test, y_test) = mnist.load_data()
# Preprocess
X_train = X_train[..., tf.newaxis] / 255.0 # shape (60000,28,28,1)
X_test = X_test[..., tf.newaxis] / 255.0
# One‑hot encode labels
y_train = tf.keras.utils.to_categorical(y_train, 10)
y_test = tf.keras.utils.to_categorical(y_test, 10)
model = models.Sequential([
layers.Conv2D(32, (3,3), activation='relu', input_shape=(28,28,1)),
layers.MaxPooling2D((2,2)),
layers.Conv2D(64, (3,3), activation='relu'),
layers.MaxPooling2D((2,2)),
layers.Flatten(),
layers.Dense(128, activation='relu'),
layers.Dropout(0.5),
layers.Dense(10, activation='softmax')
])
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
model.summary()
history = model.fit(X_train, y_train, epochs=10, batch_size=64, validation_split=0.1)
test_loss, test_acc = model.evaluate(X_test, y_test)
print('Test accuracy:', test_acc)
Visualizing filters
first_layer_weights = model.layers[0].get_weights()[0] # shape (3,3,1,32)
# Plot the first 6 filters
import matplotlib.pyplot as plt
fig, axs = plt.subplots(1,6, figsize=(15,2))
for i in range(6):
axs[i].imshow(first_layer_weights[:,:,0,i], cmap='viridis')
axs[i].axis('off')
plt.show()
Transfer learning (using a pre‑trained model)
base_model = tf.keras.applications.MobileNetV2(input_shape=(128,128,3), include_top=False, weights='imagenet')
base_model.trainable = False # freeze layers
inputs = tf.keras.Input(shape=(128,128,3))
x = base_model(inputs, training=False)
x = layers.GlobalAveragePooling2D()(x)
outputs = layers.Dense(5, activation='softmax')(x)
model_tl = tf.keras.Model(inputs, outputs)
model_tl.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
Tips
- Normalize images to
[0,1]or use ImageNet mean/std. - Use data augmentation (
tf.keras.preprocessing.image.ImageDataGenerator). - Early stopping prevents over‑training.