TensorFlow CNNs and Transfer Learning
Build convolutional networks in Keras — from scratch on image data — and fine-tune pre-trained ImageNet models for custom tasks.
Real-World Scenario
A plant disease detection app needs to classify leaf photos into 38 disease categories. Training a CNN from scratch requires millions of samples. Using transfer learning with a pre-trained EfficientNetB0 backbone, the same accuracy is achieved with just 54,000 images and 20 minutes of fine-tuning on a single GPU.
Building a CNN from Scratch
import tensorflow as tf
from tensorflow import keras
# CIFAR-10: 32x32 color images, 10 classes
(X_train, y_train), (X_test, y_test) = keras.datasets.cifar10.load_data()
# Normalize to [0, 1]
X_train = X_train.astype("float32") / 255.0
X_test = X_test.astype("float32") / 255.0
# Labels as integers (SparseCategoricalCrossentropy doesn't need one-hot)
y_train = y_train.squeeze()
y_test = y_test.squeeze()
def build_cnn(input_shape=(32, 32, 3), n_classes=10) -> keras.Model:
return keras.Sequential([
# Block 1
keras.layers.Conv2D(32, (3, 3), padding="same", activation="relu",
input_shape=input_shape),
keras.layers.BatchNormalization(),
keras.layers.Conv2D(32, (3, 3), padding="same", activation="relu"),
keras.layers.BatchNormalization(),
keras.layers.MaxPooling2D(2, 2),
keras.layers.Dropout(0.25),
# Block 2
keras.layers.Conv2D(64, (3, 3), padding="same", activation="relu"),
keras.layers.BatchNormalization(),
keras.layers.Conv2D(64, (3, 3), padding="same", activation="relu"),
keras.layers.BatchNormalization(),
keras.layers.MaxPooling2D(2, 2),
keras.layers.Dropout(0.25),
# Block 3
keras.layers.Conv2D(128, (3, 3), padding="same", activation="relu"),
keras.layers.BatchNormalization(),
keras.layers.MaxPooling2D(2, 2),
keras.layers.Dropout(0.25),
# Classifier head
keras.layers.Flatten(),
keras.layers.Dense(256, activation="relu"),
keras.layers.BatchNormalization(),
keras.layers.Dropout(0.5),
keras.layers.Dense(n_classes, activation="softmax"),
])
model = build_cnn()
model.summary()
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=1e-3),
loss="sparse_categorical_crossentropy",
metrics=["accuracy"],
)
# Data augmentation on the fly
datagen = keras.preprocessing.image.ImageDataGenerator(
rotation_range=15,
width_shift_range=0.1,
height_shift_range=0.1,
horizontal_flip=True,
zoom_range=0.1,
)
datagen.fit(X_train)
history = model.fit(
datagen.flow(X_train, y_train, batch_size=64),
epochs=30,
validation_data=(X_test, y_test),
callbacks=[
keras.callbacks.EarlyStopping(patience=5, restore_best_weights=True),
keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=3, verbose=1),
]
)
test_loss, test_acc = model.evaluate(X_test, y_test, verbose=0)
print(f"Test accuracy: {test_acc:.4f}")
Transfer Learning with a Pre-Trained Backbone
import tensorflow as tf
from tensorflow import keras
import numpy as np
# Simulate a binary classification task (cats vs dogs)
# In practice: tf.keras.utils.image_dataset_from_directory("path/to/data")
IMG_SIZE = 224
BATCH_SIZE = 32
# Create tf.data pipeline with augmentation
def make_dataset(n_samples: int, n_classes: int = 2) -> tf.data.Dataset:
rng = np.random.default_rng(42)
X = rng.integers(0, 256, (n_samples, IMG_SIZE, IMG_SIZE, 3), dtype=np.uint8)
y = rng.integers(0, n_classes, n_samples)
ds = tf.data.Dataset.from_tensor_slices((X, y))
return ds.shuffle(1000).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
train_ds = make_dataset(1000)
val_ds = make_dataset(200)
# Pre-processing and augmentation layers
augmentation = keras.Sequential([
keras.layers.RandomFlip("horizontal"),
keras.layers.RandomRotation(0.1),
keras.layers.RandomZoom(0.1),
keras.layers.RandomContrast(0.1),
])
def preprocess(image, label):
image = tf.cast(image, tf.float32)
# EfficientNet expects [0, 255] — no manual normalization needed
return image, label
def build_transfer_model(n_classes: int = 2, freeze_base: bool = True) -> keras.Model:
base_model = keras.applications.EfficientNetB0(
include_top=False, # remove the ImageNet classifier head
weights="imagenet",
input_shape=(IMG_SIZE, IMG_SIZE, 3),
)
base_model.trainable = not freeze_base
inputs = keras.Input(shape=(IMG_SIZE, IMG_SIZE, 3))
x = augmentation(inputs) # only applied during training
x = preprocess(x, None)[0]
x = base_model(x, training=False) # training=False: BatchNorm uses stored stats
x = keras.layers.GlobalAveragePooling2D()(x)
x = keras.layers.Dropout(0.3)(x)
outputs = keras.layers.Dense(n_classes, activation="softmax")(x)
return keras.Model(inputs, outputs)
# Phase 1: Train only the classification head (base frozen)
model = build_transfer_model(freeze_base=True)
model.compile(
optimizer=keras.optimizers.Adam(1e-3),
loss="sparse_categorical_crossentropy",
metrics=["accuracy"],
)
print(f"Trainable params (frozen base): {model.trainable_variables.__len__()}")
model.fit(
train_ds.map(preprocess),
validation_data=val_ds.map(preprocess),
epochs=10,
callbacks=[keras.callbacks.EarlyStopping(patience=3, restore_best_weights=True)],
)
# Phase 2: Unfreeze and fine-tune everything at a lower learning rate
def fine_tune_all(model: keras.Model, n_classes: int) -> keras.Model:
base = model.layers[4] # EfficientNetB0 layer
base.trainable = True
# Optionally freeze early layers (they learn very generic features)
for layer in base.layers[:100]:
layer.trainable = False
model.compile(
optimizer=keras.optimizers.Adam(1e-5), # 100x lower LR for fine-tuning
loss="sparse_categorical_crossentropy",
metrics=["accuracy"],
)
return model
model = fine_tune_all(model, n_classes=2)
print(f"Trainable params (unfrozen): {len(model.trainable_variables)}")
model.fit(
train_ds.map(preprocess),
validation_data=val_ds.map(preprocess),
epochs=20,
callbacks=[
keras.callbacks.EarlyStopping(patience=5, restore_best_weights=True),
keras.callbacks.ModelCheckpoint("best_model.keras", save_best_only=True),
],
)
tf.data Pipeline for Real Image Datasets
import tensorflow as tf
from pathlib import Path
# Load from a directory structure:
# data/
# train/
# cats/ *.jpg
# dogs/ *.jpg
# val/
# cats/ *.jpg
# dogs/ *.jpg
def load_image_dataset(
directory: str,
image_size: tuple = (224, 224),
batch_size: int = 32,
subset: str = "training",
) -> tf.data.Dataset:
return tf.keras.utils.image_dataset_from_directory(
directory,
labels="inferred", # class names from subdirectory names
label_mode="int",
image_size=image_size,
batch_size=batch_size,
shuffle=(subset == "training"),
seed=42,
)
def optimize_pipeline(ds: tf.data.Dataset, training: bool = True) -> tf.data.Dataset:
"""Add caching, prefetching, and normalization."""
AUTOTUNE = tf.data.AUTOTUNE
if training:
ds = ds.shuffle(1000)
ds = ds.cache() # cache after loading (speeds up epoch 2+)
ds = ds.prefetch(AUTOTUNE) # overlap data loading with GPU compute
return ds
# Custom data loading for unlabeled inference
def predict_folder(model: tf.keras.Model, image_dir: str, class_names: list[str]):
"""Run inference on all images in a directory."""
image_paths = list(Path(image_dir).glob("*.jpg"))
for path in image_paths:
img = tf.keras.utils.load_img(path, target_size=(224, 224))
arr = tf.keras.utils.img_to_array(img)
arr = tf.expand_dims(arr, 0) # add batch dimension
preds = model.predict(arr, verbose=0)
class_idx = tf.argmax(preds[0]).numpy()
confidence = preds[0][class_idx]
print(f"{path.name}: {class_names[class_idx]} ({confidence:.1%})") Frequently Asked Questions
What is transfer learning and why is it so effective?
Transfer learning reuses weights from a model trained on a large dataset (typically ImageNet with 14M images). The early layers learn universal features — edges, textures, shapes — that transfer to almost any vision task. Fine-tuning on your custom data trains only the last layers (or all layers at a lower learning rate), needing 10-100x less data than training from scratch.
When should I freeze base model layers vs fine-tune all layers?
Freeze when: your dataset is small (<1,000 images) or very similar to ImageNet. Fine-tune all layers when: you have more data (5,000+) or your domain differs significantly from ImageNet (medical images, satellite imagery). A common strategy: train with frozen base for a few epochs, then unfreeze and train at 10x lower learning rate.