Skip to main content
TensorFlow intermediate Lesson 3 of 7

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.