Aller au contenu

Reconnaître des chiffres : entraîner un CNN

Contrairement à YOLO (pré-entraîné), ici vous entraînez vous-même un réseau de neurones, de zéro, avec PyTorch. L’objectif : reconnaître des chiffres manuscrits grâce à la base MNIST, puis exposer la prédiction dans un nœud ROS 2. C’est l’occasion de voir tout le cycle de l’apprentissage : données → modèle → entraînement → inférence.

Fenêtre de terminal
pip install torch torchvision # CPU par défaut (suffisant pour MNIST)

torchvision télécharge MNIST (60 000 images d’entraînement, 10 000 de test, en 28×28 niveaux de gris) :

from torchvision import datasets
from torchvision.transforms import ToTensor
from torch.utils.data import DataLoader
train_data = datasets.MNIST(root="data", train=True, transform=ToTensor(), download=True)
test_data = datasets.MNIST(root="data", train=False, transform=ToTensor(), download=True)
loaders = {
"train": DataLoader(train_data, batch_size=100, shuffle=True, num_workers=1),
"test": DataLoader(test_data, batch_size=100, shuffle=False, num_workers=1),
}
print(train_data.data.shape) # torch.Size([60000, 28, 28])

Un petit réseau : deux couches de convolution (extraction de motifs) suivies de deux couches « pleinement connectées » (classification en 10 chiffres).

import torch.nn as nn
import torch.nn.functional as F
class CNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 10, kernel_size=5)
self.conv2 = nn.Conv2d(10, 20, kernel_size=5)
self.conv2_drop = nn.Dropout2d()
self.fc1 = nn.Linear(320, 50)
self.fc2 = nn.Linear(50, 10)
def forward(self, x):
x = F.relu(F.max_pool2d(self.conv1(x), 2))
x = F.relu(F.max_pool2d(self.conv2_drop(self.conv2(x)), 2))
x = x.view(-1, 320)
x = F.relu(self.fc1(x))
x = F.dropout(x, training=self.training)
x = self.fc2(x)
return F.log_softmax(x, dim=1)

On choisit le dispositif (GPU si dispo), l’optimiseur (Adam) et la fonction de perte (entropie croisée).

import torch
import torch.optim as optim
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CNN().to(device)
optimizer = optim.Adam(model.parameters(), lr=0.001)
loss_fn = nn.CrossEntropyLoss()

La boucle d’entraînement : pour chaque lot, on calcule la prédiction, la perte, puis on rétropropage et on met à jour les poids.

def train(epoch):
model.train()
for batch_idx, (data, target) in enumerate(loaders["train"]):
data, target = data.to(device), target.to(device)
optimizer.zero_grad()
loss = loss_fn(model(data), target)
loss.backward()
optimizer.step()
if batch_idx % 100 == 0:
print(f"Epoch {epoch} [{batch_idx * len(data)}/{len(train_data)}] perte={loss.item():.4f}")
def test():
model.eval()
correct = 0
with torch.no_grad():
for data, target in loaders["test"]:
data, target = data.to(device), target.to(device)
pred = model(data).argmax(dim=1)
correct += pred.eq(target).sum().item()
print(f"Précision test : {100 * correct / len(test_data):.1f}%")
for epoch in range(1, 6):
train(epoch)
test()
torch.save(model.state_dict(), "mnist_cnn.pth") # sauvegarde des poids

On recharge les poids et on prédit. Une image caméra doit être préparée comme MNIST : niveaux de gris, 28×28, chiffre blanc sur fond noir.

import cv2 as cv
import torch
model = CNN().to(device)
model.load_state_dict(torch.load("mnist_cnn.pth"))
model.eval()
def predire(img_bgr):
gris = cv.cvtColor(img_bgr, cv.COLOR_BGR2GRAY)
gris = cv.resize(gris, (28, 28))
gris = cv.bitwise_not(gris) # fond noir / chiffre blanc
x = torch.tensor(gris, dtype=torch.float32) / 255.0
x = x.unsqueeze(0).unsqueeze(0).to(device) # forme (1, 1, 28, 28)
with torch.no_grad():
return int(model(x).argmax(dim=1).item())

Faites apparaître un panneau-chiffre dans la simulation et testez sur une capture :

Fenêtre de terminal
ros2 launch bootcamp_vision spawn_object.launch.py object_type:=digit

Reprenez le squelette du socle commun (make_det y est défini). Le modèle est chargé une fois dans le constructeur ; à chaque image, on localise le panneau (seuillage + plus gros contour), on recadre sa boîte, on la classe, et on l’ajoute. Le class_name est le chiffre prédit ; la pose.position vient de la back-projection du centre de sa boîte.

class DigitDetector(Node):
def __init__(self):
super().__init__("digit_detector")
# ... bridge / pub (DetectionArray) / subscription comme le socle commun ...
self.model = CNN().to(device)
self.model.load_state_dict(torch.load("mnist_cnn.pth"))
self.model.eval()
def on_image(self, msg):
frame = self.bridge.imgmsg_to_cv2(msg, "bgr8")
out = DetectionArray()
out.header = msg.header
# 1) localiser le panneau : seuillage + contours
gris = cv.cvtColor(frame, cv.COLOR_BGR2GRAY)
flou = cv.GaussianBlur(gris, (5, 5), 0)
_, binaire = cv.threshold(flou, 0, 255, cv.THRESH_BINARY + cv.THRESH_OTSU)
contours, _ = cv.findContours(binaire, cv.RETR_EXTERNAL, cv.CHAIN_APPROX_SIMPLE)
contours = [c for c in contours if cv.contourArea(c) > 1000] # anti-bruit
# 2) le panneau = plus gros contour -> recadrer -> predire
if contours:
cnt = max(contours, key=cv.contourArea)
x, y, w, h = cv.boundingRect(cnt)
chiffre = predire(frame[y:y + h, x:x + w])
out.detections.append(
self.make_det(str(chiffre), chiffre, 1.0, x + w / 2, y + h / 2, w, h))
self.pub.publish(out)
Fenêtre de terminal
ros2 run mon_projet_vision digit_detector
ros2 topic echo /detections # class_name = "7", pose.position = centre back-projeté
  • Affichez la probabilité (confiance) de la prédiction en plus du chiffre.
  • Augmentez les epochs ou ajoutez une couche : la précision test bouge-t-elle ?
  • Envie d’une détection multi-objets ? Voir YOLO. Sans IA : Formes ou ArUco.

Cours conçu et animé par Etienne Schmitz