Overview¶
This lesson trains a RetinaNet object detector on a marine debris dataset using KerasCV, providing a TensorFlow/Keras counterpart to the PyTorch Faster R-CNN lesson in 3.6. Running through both lessons with the same dataset is intentional: the architectural and API differences become concrete when the data and the task are held constant.
RetinaNet as a single-stage detector. RetinaNet does not use a region proposal network. Instead, it places a fixed grid of anchor boxes across the image at multiple scales and aspect ratios, then predicts class probabilities and box offsets for every anchor in a single forward pass. This eliminates the two-stage overhead but introduces a severe class imbalance problem: for a typical image there may be 100,000 anchors, of which only a handful overlap with actual objects. Training with standard cross-entropy loss means the easy background anchors dominate the gradient, drowning out the signal from the rare positive examples.
Focal Loss. RetinaNet addresses class imbalance with Focal Loss, introduced by Lin et al. (2017). Focal Loss modifies cross-entropy by adding a modulating factor (1 - p_t)^gamma, where p_t is the model’s estimated probability for the correct class. When a sample is easy (high confidence, correct prediction), p_t is close to 1 and the factor drives the loss contribution toward 0. When a sample is hard (low confidence or wrong prediction), p_t is small and the loss is close to standard cross-entropy. With a typical gamma of 2.0, easy negatives contribute roughly 100 times less to the gradient than hard examples. This focuses training on the minority of informative samples.
KerasCV provides a high-level RetinaNet class with preset backbones, handling anchor generation, decoding, and NMS internally.
Learning Objectives¶
By the end of this lesson you will be able to:
Explain how RetinaNet differs from two-stage detectors and why Focal Loss is necessary for single-stage dense prediction.
Use KerasCV’s
from_presetAPI to load a pretrained detection model.Build a
tf.datapipeline that parses YOLO-format annotations and converts bounding box coordinates to the format KerasCV expects.Compile and train a RetinaNet model with focal classification loss and smooth L1 box regression loss.
Run inference and interpret the raw prediction output before Non-Maximum Suppression.
!pip install -q keras-cv tensorflow huggingface_hub pyyamlDataset: Win2026 TrashTank¶
The Win2026TrashTank dataset contains annotated images of marine debris items photographed in aquarium and tank settings, collected as part of a robotics competition. Annotations are in YOLO text format: each image has a corresponding .txt file where each line encodes one object as class_id cx cy w h, with all spatial values normalized to [0, 1] relative to the image dimensions.
The dataset is downloaded from Hugging Face using snapshot_download, which mirrors the entire repository to a local directory and handles incremental re-downloads. If the repository is private, a Hugging Face access token stored via userdata.get('HF_TOKEN') in Colab is required.
This is the same dataset used in lesson 3.6 (Faster R-CNN with PyTorch). Using an identical dataset across both lessons allows a direct comparison between the PyTorch and TensorFlow/Keras workflows: the data loading, bounding box format conversion, training loop, and evaluation patterns all differ, but the underlying detection task and the annotations are the same.
import os
from huggingface_hub import snapshot_download
# In Colab, use userdata.get('HF_TOKEN').
import os
token = os.environ.get('HF_TOKEN', None)
print('Downloading Win2026TrashTank dataset...')
dataset_path = snapshot_download(
repo_id='OceanCV/Win2026TrashTank',
repo_type='dataset',
local_dir='./Dataset',
token=token
)
print(f"Dataset ready at: {dataset_path}")Building a TensorFlow Data Pipeline¶
KerasCV’s RetinaNet expects bounding boxes in xyxy format: absolute pixel coordinates [x_min, y_min, x_max, y_max]. YOLO’s format is normalized [cx, cy, w, h]. The parse_yolo_txt function performs this conversion by multiplying the center and size values by the image width and height, then computing the corner coordinates:
x_min = (cx - w/2) * image_width
y_min = (cy - h/2) * image_height
x_max = (cx + w/2) * image_width
y_max = (cy + h/2) * image_heightThe tf.data.Dataset pipeline applies this parsing function via .map(), then shuffles, batches, and prefetches. Shuffling before batching ensures that each batch contains a random mix of scenes rather than all images from the same directory or sequence. Prefetching overlaps data loading with GPU computation, preventing the GPU from sitting idle while the CPU loads the next batch from disk.
The ragged=True option in the batching step is necessary because different images contain different numbers of annotated objects. A standard dense batch requires all elements along every dimension to have the same size, which would force you to pad every image’s annotation list to the length of the longest one in the batch. Ragged tensors avoid this padding, storing variable-length sequences efficiently. KerasCV’s training pipeline handles ragged bounding box tensors natively.
import tensorflow as tf
import keras_cv
import yaml
import glob
import numpy as np
# Load classes
with open('./Dataset/data.yaml', 'r') as f:
data_yaml = yaml.safe_load(f)
classes = data_yaml['names']
num_classes = len(classes)
print('Classes:', classes)
# Parse YOLO txt
def parse_yolo_txt(txt_path, img_w, img_h):
boxes, class_ids = [], []
if os.path.exists(txt_path):
with open(txt_path, 'r') as f:
for line in f:
parts = line.strip().split()
if len(parts) >= 5:
cls_id = int(parts[0])
cx, cy, w, h = map(float, parts[1:5])
# Convert to absolute xyxy
xmin = (cx - w / 2) * img_w
ymin = (cy - h / 2) * img_h
xmax = (cx + w / 2) * img_w
ymax = (cy + h / 2) * img_h
boxes.append([xmin, ymin, xmax, ymax])
class_ids.append(cls_id)
return np.array(boxes, dtype=np.float32), np.array(class_ids, dtype=np.int32)
def load_image_and_labels(img_path):
# Read image
img = tf.io.read_file(img_path)
img = tf.image.decode_jpeg(img, channels=3)
img_shape = tf.shape(img)
img_w, img_h = tf.cast(img_shape[1], tf.float32), tf.cast(img_shape[0], tf.float32)
# Corresponding txt
txt_path = tf.strings.regex_replace(img_path, '.jpg', '.txt')
txt_path = tf.strings.regex_replace(txt_path, 'images', 'labels')
def get_boxes(p, w, h):
boxes, classes = parse_yolo_txt(p.numpy().decode('utf-8'), w.numpy(), h.numpy())
if len(boxes) == 0:
boxes = np.zeros((0, 4), dtype=np.float32)
classes = np.zeros((0,), dtype=np.int32)
return boxes, classes
boxes, class_ids = tf.py_function(get_boxes, [txt_path, img_w, img_h], [tf.float32, tf.int32])
boxes.set_shape([None, 4])
class_ids.set_shape([None])
bounding_boxes = {"classes": class_ids, "boxes": boxes}
return img, bounding_boxes
def create_dataset(image_dir, batch_size=4):
paths = glob.glob(os.path.join(image_dir, '*.jpg'))
ds = tf.data.Dataset.from_tensor_slices(paths)
ds = ds.map(load_image_and_labels, num_parallel_calls=tf.data.AUTOTUNE)
# Resize via KerasCV (handles bbox resizing automatically)
resizer = keras_cv.layers.Resizing(640, 640, bounding_box_format="xyxy")
def apply_resizer(img, bboxes):
inputs = {"images": img, "bounding_boxes": bboxes}
outputs = resizer(inputs)
return outputs["images"], outputs["bounding_boxes"]
ds = ds.map(apply_resizer, num_parallel_calls=tf.data.AUTOTUNE)
# Pad batches (required for varying number of boxes per image)
ds = ds.padded_batch(batch_size, padding_values=(0.0, {"classes": -1, "boxes": 0.0}))
return ds.prefetch(tf.data.AUTOTUNE)
train_ds = create_dataset('./Dataset/train/images', batch_size=4)
val_ds = create_dataset('./Dataset/val/images', batch_size=4)
print("Datasets created successfully.")Monitoring with TensorBoard¶
TensorBoard logs training metrics so you can track progress across epochs without parsing console output. In this lesson, Keras logs metrics automatically when you pass a TensorBoard callback to model.fit. The callback writes epoch-level summaries (loss, learning rate, and any compiled metrics) to the directory specified in --logdir.
The cell above uses the magic command %tensorboard --logdir ./runs/keras_detect to launch TensorBoard embedded in the Colab notebook output. Run this cell before the training cell so the dashboard is ready when the first epoch completes.
The primary metric to watch is total loss, which is the sum of the focal classification loss and the smooth L1 box regression loss. Keras also reports these as separate keys (classification_loss and box_loss) if the model is compiled with named loss objects. A diverging total loss, where it rises rather than falls, usually indicates a learning rate that is too high. A loss that drops sharply in epoch 1 and then plateaus is normal behavior for fine-tuning a pretrained backbone: the backbone weights are already useful, so the improvement from random initialization is fast and the remaining gain requires careful tuning.
%load_ext tensorboard
%tensorboard --logdir ./runs/keras_detectTraining RetinaNet¶
keras_cv.models.RetinaNet.from_preset('resnet50_v2_imagenet') loads a RetinaNet model with a ResNet50-V2 backbone pretrained on ImageNet. ResNet50-V2 uses pre-activation residual blocks (batch norm and activation before the convolution rather than after), which typically provides slightly better accuracy than the original ResNet50 at the same computational cost.
The model is compiled with two loss functions:
Focal Loss for classification. The
gammaparameter (default 2.0) controls how aggressively easy examples are down-weighted. Higher values focus training more sharply on hard examples.alpha(default 0.25) balances the contribution of foreground and background classes at the loss level.Smooth L1 loss (also called Huber loss) for box regression. This behaves like L2 loss for small errors and L1 loss for large errors, making it less sensitive to outlier annotations than pure L2.
bounding_box_format='xyxy' must match the format used in the data pipeline. KerasCV uses this argument internally when computing IoU between predicted and ground-truth boxes during loss calculation. A mismatch between the declared format and the actual tensor layout will produce wrong IoU values and a loss that does not converge.
This demo trains for 1 epoch due to resource constraints. One epoch is not enough to produce a useful detector, but it confirms that the entire pipeline, from data loading through gradient updates, runs without errors. For real training, increase epochs to at least 20 and consider using a learning rate schedule with warm-up.
import keras
# Note: Due to resource limits, we train for 1 epoch. Increase this for real training.
epochs = 1
model = keras_cv.models.RetinaNet.from_preset(
"resnet50_v2_imagenet",
num_classes=num_classes,
bounding_box_format="xyxy"
)
# Prepare optimizer and metrics
base_lr = 0.005
optimizer = keras.optimizers.SGD(learning_rate=base_lr, momentum=0.9)
model.compile(
classification_loss="focal",
box_loss="smoothl1",
optimizer=optimizer,
)
callbacks = [
keras.callbacks.TensorBoard(log_dir='./runs/keras_detect'),
keras.callbacks.EarlyStopping(patience=5)
]
model.fit(
train_ds,
validation_data=val_ds,
epochs=epochs,
callbacks=callbacks
)
print("Training complete.")Running Inference¶
model.predict(images) returns a dict with boxes and confidence keys. Unlike the PyTorch Faster R-CNN output, which returns post-NMS detections, the raw KerasCV output includes predictions for all anchors before Non-Maximum Suppression filtering. The shape of predictions['boxes'] will therefore be large: one row per anchor per image in the batch.
To obtain final detections with NMS applied, use model.decode_predictions(pred, images). This method applies the confidence threshold and NMS to reduce the thousands of raw anchor predictions to a small set of non-overlapping boxes. The output format from decode_predictions is a list of dicts with keys boxes, confidence, and classes, one dict per image in the batch.
In a production evaluation you would loop over the validation dataset, collect predicted and ground-truth boxes for each image, and compute mAP using a metric library. The keras_cv.metrics.COCOMeanAveragePrecision metric handles this accumulation, including the COCO standard IoU thresholds from 0.50 to 0.95 in steps of 0.05 and the small/medium/large area breakdowns. The torchmetrics approach used in lesson 3.6 produces comparable numbers across frameworks.
# Run inference on a batch of validation images
for images, boxes in val_ds.take(1):
preds = model.predict(images)
print("Predictions shape:", preds['boxes'].shape)
breakReflecting on Results¶
Consider the following questions as you review your training output and compare this lesson with lesson 3.6:
Framework comparison. The Keras RetinaNet and PyTorch Faster R-CNN lessons used identical data but different frameworks, architectures, and APIs. What were the key differences in how you formatted the bounding boxes, how the loss was computed, and how the training loop was expressed? Which framework felt more explicit about what was happening inside the model, and which felt more abstract?
Focal Loss and marine debris. Marine debris detection involves a strong foreground/background imbalance: most of the pixels in any given image are water, substrate, or aquatic life rather than debris. Explain in your own words why Focal Loss addresses this imbalance more effectively than standard cross-entropy. What value of
gammawould you try first if your model was confidently predicting background for every anchor, and why?Edge deployment. Suppose you want to run this RetinaNet on an embedded system attached to an AUV (Autonomous Underwater Vehicle) with a Coral TPU or NVIDIA Jetson. Outline the steps you would take to convert the trained KerasCV model to TensorFlow Lite, including quantization, and describe one accuracy trade-off you would expect to accept in exchange for reduced latency and power consumption.