Entrenamiento de un modelo de detección de imágenes de RetinaNet

Entrene un modelo de detección de objetos RetinaNet desde cero en tiempo de ejecución de IA mediante PyTorch y torchvision. RetinaNet es un modelo de detección de objetos de una sola fase que usa una red piramidal de características (Red Piramidal de Características, FPN) y una pérdida focal para gestionar el desequilibrio de clases.

En el cuaderno se describe lo siguiente:

  • Carga y transformación del conjunto de datos coco para la detección de objetos
  • Entrenamiento de un modelo RetinaNet con red troncal ResNet-50 en una sola GPU
  • Escalado del entrenamiento en varias GPU mediante Distributed Data Parallel (DDP)
  • Registro de métricas de entrenamiento con MLflow

Nota

Este ejemplo requiere el entorno de IA Databricks versión 5 o superior.

Conectar al cómputo de GPU sin servidor

Para ejecutar este notebook, conéctese a un cómputo de GPU sin servidor con 1xA10 para el entrenamiento con una sola GPU o 8xH100 para la sección de entrenamiento distribuido.

  1. Haga clic en el selector de proceso del cuaderno en la parte superior derecha y seleccione GPU sin servidor.
  2. A la derecha, haga clic en el botón de entorno.
  3. Seleccione 1xA10 o 8xH100 como Acelerador.
  4. Seleccione AI v5 como entorno y haga clic en Aplicar.

Instalación de paquetes necesarios

Instale pycocotools para utilidades de conjuntos de datos coco y reinicie el entorno de Python para cargar el nuevo paquete.

%pip install pycocotools
dbutils.library.restartPython()

Configurar rutas de acceso del catálogo de Unity utilizando widgets

Defina widgets para especificar el catálogo de Unity Catalog, el esquema y el volumen donde se almacena el conjunto de datos COCO.

dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_volume", "coco_data")

UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_VOLUME = dbutils.widgets.get("uc_volume")

print(f"UC_CATALOG: {UC_CATALOG}")
print(f"UC_SCHEMA: {UC_SCHEMA}")
print(f"UC_VOLUME: {UC_VOLUME}")

Importación de bibliotecas de PyTorch

Importar torch y torchvision para crear y entrenar el modelo de detección de imágenes.

import torch
import torchvision

Importar clases de modelo y conjunto de datos

Importe la arquitectura del modelo RetinaNet y las utilidades del conjunto de datos COCO desde torchvision.

import os
from torchvision.models.detection import retinanet_resnet50_fpn_v2

# For this example we will be using a default Dataset from torch
from torchvision.datasets import CocoDetection

Importación de utilidades de entrenamiento distribuido

Importe los módulos de entrenamiento distribuido de PyTorch y el decorador de GPU distribuido y sin servidor para el entrenamiento con múltiples GPU.

Definición de hiperparámetros de entrenamiento y rutas de acceso de datos

Configure las rutas de acceso de datos, el tamaño del lote, el número de clases, la velocidad de aprendizaje y otros parámetros de entrenamiento. Ajuste BATCH_SIZE y NUM_EPOCHS en función del tipo de GPU y los requisitos de entrenamiento.

import torch.distributed as dist
from serverless_gpu import distributed

DATA_PATH = f"/Volumes/{UC_CATALOG}/{UC_SCHEMA}/{UC_VOLUME}/"
TRAIN_IMG_PATH = os.path.join(DATA_PATH, "val2017")
TRAIN_ANN_PATH = os.path.join(DATA_PATH, "annotations", "instances_val2017.json")

BATCH_SIZE = 2 # Please use batch size of 8 with H100 for best performance
NUM_CLASSES = 91
LEARNING_RATE = 0.005
MOMENTUM = 0.9
WEIGHT_DECAY = 0.0005
NUM_EPOCHS = 1 # Update num_epochs accordingly

Inicialización del modelo RetinaNet

Crea un modelo RetinaNet con una estructura ResNet-50 sin pesos preentrenados, configurado para el número de clases del conjunto de datos COCO.

# Since we are training the model from scratch, we need to initialize weights to None
model = retinanet_resnet50_fpn_v2(weights=None, num_classes=NUM_CLASSES)

Transformación de imágenes y anotaciones en entradas de modelo

El modelo requiere entradas como tensores con forma (C, H, W), tipo de datos float32 y rango normalizado (de 0,0 a 1,0). La función get_transform convierte imágenes de PIL y aplica técnicas de aumento de datos. La CocoWrapper clase envuelve el conjunto de datos COCO para formatear los cuadros delimitadores y las etiquetas correctamente.

from torchvision.transforms import v2
from torchvision import tv_tensors

def get_transform(train):
  transforms = []
  transforms.append(v2.ToImage())
  transforms.append(v2.ToDtype(torch.float32, scale=True))
  if train:
    transforms.append(v2.RandomHorizontalFlip())
  return v2.Compose(transforms)

class CocoWrapper(CocoDetection):
  def __init__(self, root, annFile, transforms=None):
    super().__init__(root, annFile)
    self._transforms = transforms

  def __getitem__(self, idx):
    img, target = super().__getitem__(idx)
    image_id = self.ids[idx]

    boxes = []
    labels = []

    for obj in target:
      x, y, w, h = obj["bbox"]

      boxes.append([x, y, x + w, y + h])
      labels.append(obj["category_id"])

    if len(boxes) == 0:
      boxes = torch.zeros((0, 4), dtype=torch.float32)
      labels = torch.zeros((0,), dtype=torch.int64)
    else:
      boxes = torch.as_tensor(boxes, dtype=torch.float32)
      labels = torch.as_tensor(labels, dtype=torch.int64)

    w, h = img.size
    boxes = torchvision.tv_tensors.BoundingBoxes(
      data=boxes,
      format=torchvision.tv_tensors.BoundingBoxFormat.XYXY,
      canvas_size=(h, w)
    )

    final_target = {
      "boxes": boxes,
      "labels": labels,
      "image_id": torch.tensor([image_id])
    }

    if self._transforms is not None:
      img, final_target = self._transforms(img, final_target)

    return img, final_target

dataset = CocoWrapper(
  root = TRAIN_IMG_PATH,
  annFile=TRAIN_ANN_PATH,
  transforms=get_transform(train=True)
)

# Sanity Check
img, target = dataset[0]

print("Image type:", type(img))
print("Image shape:", img.shape)        # should be [3, H, W]
print("Image dtype:", img.dtype)

print("\nTarget keys:", target.keys())
print("Boxes shape:", target["boxes"].shape)
print("Labels shape:", target["labels"].shape)
print("Image ID:", target["image_id"])

Creación del cargador de datos

Defina un DataLoader con agrupamiento personalizado para procesar por lotes imágenes y objetivos para el entrenamiento.

from torch.utils.data import DataLoader

def collate_fn(batch):
    images, targets = list(zip(*batch))
    return list(images), list(targets)

train_loader = DataLoader(
    dataset,
    batch_size=BATCH_SIZE,
    shuffle=True,
    num_workers=16,
    collate_fn=collate_fn,
    pin_memory=True,
    prefetch_factor=2 # Please use a prefetch_factor of 4 with H100 for best performance
)

Configuración del optimizador

Configure el optimizador SGD con parámetros de velocidad de aprendizaje, impulso y descomposición de peso.

device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')
model.to(device)

params = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.SGD(
    params,
    lr=LEARNING_RATE,
    momentum=MOMENTUM,
    weight_decay=WEIGHT_DECAY
)

Entrenamiento del modelo en una sola GPU

Ejecuta el bucle de entrenamiento durante el número especificado de épocas, registrando las métricas de pérdida en MLflow.

import time
import mlflow

model.train()

lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)

with mlflow.start_run():
    for epoch in range(NUM_EPOCHS):
        start_time = time.time()
        epoch_loss = 0

        for i, (images, targets) in enumerate(train_loader):

            images = list(image.to(device) for image in images)
            targets = [{k: v.to(device) for k, v in t.items()} for t in targets]
            loss_dict = model(images, targets)
            losses = sum(loss for loss in loss_dict.values())
            optimizer.zero_grad()
            losses.backward()
            optimizer.step()

            epoch_loss += losses.item()
            mlflow.log_metric("loss", losses.item(), step=epoch * len(train_loader) + i)

            if i % 50 == 0:
                print(f"Epoch {epoch+1} | Step {i}/{len(train_loader)} | Loss: {losses.item():.4f}")

        lr_scheduler.step()

        end_time = time.time()
        avg_loss = epoch_loss / len(train_loader)
        mlflow.log_metric("epoch_avg_loss", avg_loss, step=epoch)
        print(f"Epoch {epoch+1} Finished! Avg Loss: {avg_loss:.4f} | Time: {(end_time - start_time)/60:.2f} min")

print("Training Complete.")

Entrena con paralelismo de datos distribuido (DDP)

Escalar el entrenamiento entre varias GPUs usando el decorador @distributed. Este enfoque lee el conjunto de datos directamente desde el volumen del catálogo de Unity y usa DistributedSampler para particionar el conjunto de datos entre GPU. Todas las importaciones, transformaciones, el conjunto de datos y dataLoader se vuelven a definir dentro de la función de entrenamiento, según lo requiera el modelo de ejecución distribuido.

from datetime import timedelta

BATCH_SIZE_PER_GPU = 8  # for better performance with H100

@distributed(gpus=8, gpu_type='H100')
def train_distributed():
    import os
    import torch
    import torch.distributed as dist
    import time
    import torchvision
    import mlflow
    from torch.nn.parallel import DistributedDataParallel as DDP
    from torch.utils.data import DataLoader
    from torch.utils.data.distributed import DistributedSampler
    from torchvision.models.detection import retinanet_resnet50_fpn_v2
    from torchvision.transforms import v2
    from torchvision import tv_tensors
    from torchvision.datasets import CocoDetection

    def get_transform(train):
        transforms = []
        transforms.append(v2.ToImage())
        transforms.append(v2.ToDtype(torch.float32, scale=True))
        if train:
            transforms.append(v2.RandomHorizontalFlip())
        return v2.Compose(transforms)

    def collate_fn(batch):
        images, targets = list(zip(*batch))
        return list(images), list(targets)

    class CocoWrapper(CocoDetection):
        def __init__(self, root, annFile, transforms=None):
            super().__init__(root, annFile)
            self._transforms = transforms

        def __getitem__(self, idx):
            img, target = super().__getitem__(idx)
            image_id = self.ids[idx]

            boxes = []
            labels = []

            for obj in target:
                x, y, w, h = obj["bbox"]

                boxes.append([x, y, x + w, y + h])
                labels.append(obj["category_id"])

            if len(boxes) == 0:
                boxes = torch.zeros((0, 4), dtype=torch.float32)
                labels = torch.zeros((0,), dtype=torch.int64)
            else:
                boxes = torch.as_tensor(boxes, dtype=torch.float32)
                labels = torch.as_tensor(labels, dtype=torch.int64)

            w, h = img.size
            boxes = torchvision.tv_tensors.BoundingBoxes(
                data=boxes,
                format=torchvision.tv_tensors.BoundingBoxFormat.XYXY,
                canvas_size=(h, w)
            )

            final_target = {
                "boxes": boxes,
                "labels": labels,
                "image_id": torch.tensor([image_id])
            }

            if self._transforms is not None:
                img, final_target = self._transforms(img, final_target)

            return img, final_target

    dist.init_process_group(backend="nccl", timeout=timedelta(minutes=30))

    rank = int(os.environ["RANK"])
    local_rank = int(os.environ["LOCAL_RANK"])
    world_size = int(os.environ["WORLD_SIZE"])

    torch.cuda.set_device(local_rank)
    device = torch.device(f"cuda:{local_rank}")

    # Read the dataset directly from the Unity Catalog volume on every rank.
    train_img_path = os.path.join(DATA_PATH, "val2017")
    train_ann_path = os.path.join(DATA_PATH, "annotations", "instances_val2017.json")

    dataset = CocoWrapper(
        root=train_img_path,
        annFile=train_ann_path,
        transforms=get_transform(train=True)
    )

    train_sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank, shuffle=True)

    train_loader = DataLoader(
        dataset,
        batch_size=BATCH_SIZE_PER_GPU,
        shuffle=False,
        num_workers=8, # 8 workers * 8 GPUs = 64 total CPU threads
        collate_fn=collate_fn,
        pin_memory=True,
        prefetch_factor=4,
        sampler=train_sampler
    )

    model = retinanet_resnet50_fpn_v2(weights=None, num_classes=NUM_CLASSES)
    model.to(device)

    model = DDP(model, device_ids=[local_rank])

    params = [p for p in model.parameters() if p.requires_grad]
    optimizer = torch.optim.SGD(params, lr=0.04, momentum=0.9, weight_decay=0.0005)
    lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)

    model.train()
    if rank == 0:
        print(f"Training on {world_size} GPUs. Global Batch Size: {BATCH_SIZE_PER_GPU * world_size}")

    with mlflow.start_run():
        for epoch in range(NUM_EPOCHS):
            train_sampler.set_epoch(epoch)

            start_time = time.time()
            epoch_loss = 0

            for i, (images, targets) in enumerate(train_loader):
                images = [image.to(device) for image in images]
                targets = [{k: v.to(device) for k, v in t.items()} for t in targets]

                loss_dict = model(images, targets)
                losses = sum(loss for loss in loss_dict.values())

                optimizer.zero_grad()
                losses.backward()
                optimizer.step()

                epoch_loss += losses.item()

                if rank == 0:
                    mlflow.log_metric("loss", losses.item(), step=epoch * len(train_loader) + i)

                if rank == 0 and i % 50 == 0:
                    print(f"Rank 0 | Step {i}/{len(train_loader)} | Loss: {losses.item():.4f}")

            lr_scheduler.step()

            if rank == 0:
                avg_loss = epoch_loss / len(train_loader)
                mlflow.log_metric("epoch_avg_loss", avg_loss, step=epoch)
                print(f"Epoch {epoch+1} Finished! Avg Loss: {avg_loss:.4f} | Time: {(time.time() - start_time)/60:.2f} min")

    dist.destroy_process_group()

train_distributed.distributed()

Pasos siguientes

Cuaderno de ejemplo

Entrenamiento de un modelo de detección de imágenes de RetinaNet

Obtención del cuaderno