Entraîner un modèle de détection d’image RetinaNet

Entraîner un modèle de détection d’objet RetinaNet à partir de zéro sur AI Runtime à l’aide de PyTorch et de torchvision. RetinaNet est un modèle de détection d’objet à phase unique qui utilise un FPN (Feature Pyramid Network) et une perte focale pour gérer le déséquilibre des classes.

Le bloc-notes couvre les points suivants :

  • Chargement et transformation du jeu de données COCO pour la détection d’objets
  • Entraînement d'un modèle RetinaNet avec l'architecture ResNet-50 sur un seul GPU.
  • Mise à l’échelle de l’entraînement sur plusieurs GPU à l’aide de Distributed Data Parallel (DDP)
  • Journalisation des métriques d’apprentissage avec MLflow

Note

Cet exemple nécessite l’environnement IA Databricks version 5 ou supérieure.

Connecter au calcul GPU sans serveur

Pour exécuter ce notebook, connectez-vous à des ressources de calcul GPU serverless avec 1xA10 pour un entraînement sur un seul GPU ou 8xH100 pour la section distribuée.

  1. Cliquez sur le sélecteur de calcul du notebook en haut à droite et sélectionnez GPU serverless.
  2. Sur le côté droit, cliquez sur le bouton Environnement.
  3. Sélectionnez 1xA10 ou 8xH100 comme accélérateur.
  4. Sélectionnez AI v5 comme environnement, puis cliquez sur Appliquer.

Installer les packages requis

Installez pycocotools pour les utilitaires de jeu de données COCO et redémarrez l’environnement Python pour charger le nouveau package.

%pip install pycocotools
dbutils.library.restartPython()

Configurer des chemins de catalogue Unity avec des widgets

Définissez des widgets pour spécifier le catalogue, le schéma et le volume du catalogue Unity où le jeu de données COCO est stocké.

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}")

Importer des bibliothèques PyTorch

Importez torch et torchvision pour générer et entraîner le modèle de détection d’image.

import torch
import torchvision

Importer des classes de modèle et de jeu de données

Importez l’architecture du modèle RetinaNet et les utilitaires de jeu de données COCO à partir de 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

Importer des utilitaires d’entraînement distribués

Importez les modules d'entraînement distribués de PyTorch, ainsi que le décorateur distribué GPU sans serveur pour l'entraînement sur plusieurs GPU.

Définir des hyperparamètres d’entraînement et des chemins de données

Configurez les chemins de données, la taille du lot, le nombre de classes, le taux d’apprentissage et d’autres paramètres d’apprentissage. Ajustez BATCH_SIZE et NUM_EPOCHS en fonction des besoins en matière de formation et de type GPU.

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

Initialiser le modèle RetinaNet

Créez un modèle RetinaNet avec l’architecture de base ResNet-50 sans poids préformés, configuré pour le nombre de classes du jeu de données 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)

Transformer des images et des annotations en entrées de modèle

Le modèle nécessite des entrées en tant que tenseurs avec forme (C, H, W), type de données float32 et plage normalisée (0,0 à 1,0). La get_transform fonction convertit les images PIL et applique l’augmentation des données. La classe CocoWrapper encapsule le jeu de données COCO afin de formater correctement les cadres englobants et les étiquettes.

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"])

Créer le chargeur de données

Définissez un DataLoader avec une collation personnalisée pour regrouper les images et les cibles en lots pour la formation.

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
)

Configurer l’optimiseur

Configurez l’optimiseur SGD avec le taux d’apprentissage, l’élan et les paramètres de dégradation du poids.

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
)

Entraîner le modèle sur un seul GPU

Exécuter la boucle d'entraînement pour le nombre spécifié d'épochs, consigner les métriques de perte dans 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.")

Entraînement avec DDP (parallélisme distribué des données)

Mettez à l’échelle l’entraînement sur plusieurs GPU à l’aide de l’élément décoratif @distributed. Cette approche lit le jeu de données directement à partir du volume catalogue Unity et utilise DistributedSampler pour partitionner le jeu de données entre les GPU. Toutes les importations, transformations, jeux de données et DataLoader sont redéfinis à l’intérieur de la fonction d’entraînement, comme requis par le modèle d’exécution distribué.

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()

Étapes suivantes

Exemple de notebook

Entraîner un modèle de détection d’image RetinaNet

Obtenir un ordinateur portable