Aller au contenu principal

Module 10 — Mise en service d'un modèle ONNX

Le fichier .onnx optimisé, quantifié, vérifié, mesuré n'a de valeur qu'à travers un service qui l'expose à un client. Ce module conclut le fil rouge en montrant à quoi ressemble une API d'inférence minimale, robuste et raisonnablement rapide, avec ONNX Runtime derrière FastAPI. On termine par un aperçu d'ONNX Runtime Web qui permet d'exécuter le même modèle côté navigateur.

Une seule session, partagée

La règle de base : une session ONNX Runtime est chère à créer, bon marché à appeler. La créer à chaque requête gaspille des centaines de millisecondes en initialisation, allocation d'arène et compilation potentielle de kernels. Le bon design instancie une session au démarrage du service et la partage entre requêtes.

Sur CPU, une session peut être appelée par plusieurs fils en même temps ; ONNX Runtime sérialise en interne l'accès aux kernels. Sur GPU, l'accès à une session unique reste correct : les requêtes se mettent en file d'attente naturellement au niveau de CUDA. Il n'est pas nécessaire — et il est même contre-productif — de créer une session par requête ou par fil.

# service.py
from contextlib import asynccontextmanager
from fastapi import FastAPI, UploadFile
from io import BytesIO
from PIL import Image
import numpy as np
import onnxruntime as ort

CLASSES = ["T-shirt", "Pantalon", "Pull", "Robe", "Manteau",
"Sandale", "Chemise", "Basket", "Sac", "Bottine"]

MOYENNE = np.array([0.485, 0.456, 0.406], dtype=np.float32).reshape(1, 3, 1, 1)
ECART = np.array([0.229, 0.224, 0.225], dtype=np.float32).reshape(1, 3, 1, 1)

def creer_session():
options = ort.SessionOptions()
options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
options.intra_op_num_threads = 4
options.inter_op_num_threads = 1

session = ort.InferenceSession(
"resnet18_fashion_int8_dyn.onnx",
sess_options=options,
providers=["CPUExecutionProvider"],
)
actifs = session.get_providers()
if actifs != ["CPUExecutionProvider"]:
raise RuntimeError(f"Fournisseurs inattendus : {actifs}")
return session

@asynccontextmanager
async def cycle_de_vie(app: FastAPI):
app.state.session = creer_session()
yield
del app.state.session

app = FastAPI(lifespan=cycle_de_vie)

Ce squelette respecte trois principes :

  • La session est créée une fois au démarrage, via le lifespan de FastAPI.
  • Le fournisseur est vérifié avec get_providers() ; un repli silencieux sur CPU aurait été rattrapé si on visait CUDA.
  • Les paramètres de fils sont explicites : 4 fils intra-opérateur, 1 fil inter-opérateur, ce qui convient à un serveur qui reçoit plusieurs requêtes concurrentes.

Le prétraitement doit accompagner le modèle

Le point le plus souvent oublié en production : la transformation d'entrée. Le modèle a été entraîné sur des images normalisées avec une moyenne et un écart-type précis. Si le service reçoit des images brutes non normalisées, le modèle rend des prédictions plausibles mais fausses — pas d'erreur, pas d'exception, juste des probabilités mal calibrées.

def pretraiter(image_pil: Image.Image) -> np.ndarray:
"""Reproduit exactement le prétraitement du module 4 du cours PyTorch."""
# 1. Convertir en RGB (ResNet18 attend 3 canaux même si Fashion-MNIST est en gris)
image_pil = image_pil.convert("RGB")

# 2. Redimensionner à 256 puis recadrer au centre à 224
image_pil = image_pil.resize((256, 256), Image.BILINEAR)
marge = (256 - 224) // 2
image_pil = image_pil.crop((marge, marge, marge + 224, marge + 224))

# 3. Vers un tableau numpy (H, W, C) puis (1, C, H, W) et float32 dans [0, 1]
tableau = np.asarray(image_pil, dtype=np.float32) / 255.0
tableau = tableau.transpose(2, 0, 1)[None, :, :, :]

# 4. Normalisation ImageNet, identique à celle de l'entraînement
tableau = (tableau - MOYENNE) / ECART
return tableau.astype(np.float32)

Ce code n'est pas fantaisiste : chaque ligne reproduit une étape du pipeline d'entraînement. Toute divergence — même l'ordre RGB au lieu de BGR — dégrade la précision de plusieurs points. La façon la plus sûre d'éviter cette divergence : coller ce prétraitement dans un fichier pretraitement.py versionné à côté du fichier .onnx, en documentant qu'il correspond au commit qui a produit le modèle.

Le point d'entrée HTTP

@app.post("/predire")
async def predire(fichier: UploadFile):
contenu = await fichier.read()
image = Image.open(BytesIO(contenu))
entree = pretraiter(image)

logits = app.state.session.run(None, {"entree": entree})[0]
probas = softmax(logits[0])
indice = int(probas.argmax())

return {
"classe": CLASSES[indice],
"indice": indice,
"probabilite": float(probas[indice]),
}

def softmax(x: np.ndarray) -> np.ndarray:
x = x - x.max()
exp_x = np.exp(x)
return exp_x / exp_x.sum()

Trois observations.

Le softmax est côté service. Le modèle exporté ne l'applique pas (bonne pratique : garder le modèle en logits, ce qui préserve la précision numérique et laisse au service la responsabilité de la calibration).

Aucune dépendance à PyTorch ou TensorFlow. Le service n'importe que onnxruntime, numpy, PIL et fastapi. L'image Docker fait 400 Mio au lieu de 3 Gio.

La sortie est un dictionnaire simple, sérialisable en JSON par FastAPI sans effort.

Servir un lot pour tirer parti du modèle

Servir une image à la fois sur GPU sous-utilise le matériel (module 8). Le service peut regrouper les requêtes avec une petite file d'attente : accumuler jusqu'à 16 requêtes ou 5 millisecondes, puis appeler la session avec un lot.

import asyncio
from asyncio import Queue

class Regroupeur:
def __init__(self, session, taille_max=16, delai_ms=5):
self.session = session
self.taille_max = taille_max
self.delai_ms = delai_ms
self.file: Queue = Queue()
self._tache = None

async def demarrer(self):
self._tache = asyncio.create_task(self._boucle())

async def _boucle(self):
while True:
requetes = [await self.file.get()]
t_debut = asyncio.get_event_loop().time()
while (len(requetes) < self.taille_max
and (asyncio.get_event_loop().time() - t_debut) * 1000 < self.delai_ms):
try:
requetes.append(await asyncio.wait_for(self.file.get(),
timeout=self.delai_ms / 1000))
except asyncio.TimeoutError:
break

lot = np.concatenate([r["entree"] for r in requetes], axis=0)
logits = self.session.run(None, {"entree": lot})[0]
for r, l in zip(requetes, logits):
r["future"].set_result(l)

Ce composant reste pédagogique — un vrai service en production utiliserait Triton de NVIDIA ou une bibliothèque de mise en lot plus robuste. Mais le principe est là : les requêtes arrivent une à une, sont regroupées en lot, et chaque cliente reçoit sa réponse individuelle.

Vérifier le service en production

Trois vérifications à automatiser au déploiement.

  • Un test de bout en bout : le service reçoit une image connue, renvoie la classe attendue. Le nom du test s'appelle « détection de désalignement » : si le modèle prédisait « Basket » sur cette image en test hors ligne, il doit prédire « Basket » en production.
  • Une mesure de latence à chaud : ping le service avec 500 images en série, vérifier que p95 reste sous le budget alloué (par exemple 20 ms). Un p95 qui dérive signale une régression matérielle (nouveau serveur sans VNNI par exemple).
  • Une somme de contrôle du fichier .onnx : loguer le SHA-256 du modèle chargé, comparer à celui commité. Un modèle remplacé silencieusement lors d'un docker pull est un incident qu'on n'oublie pas.

ONNX Runtime Web : le modèle dans le navigateur

Une propriété unique d'ONNX est qu'il tourne aussi dans un navigateur via WebAssembly ou WebGPU. La bibliothèque onnxruntime-web charge le même fichier .onnx et l'exécute côté client, sans jamais envoyer les données à un serveur.

import * as ort from "onnxruntime-web";

async function predire(imageArray) {
const session = await ort.InferenceSession.create(
"/models/resnet18_fashion.onnx",
{ executionProviders: ["wasm"] },
);
const entree = new ort.Tensor("float32", imageArray, [1, 3, 224, 224]);
const sortie = await session.run({ entree });
return sortie.logits.data;
}

Deux cas d'usage typiques.

  • Une application où la latence réseau est intolérable : classification d'image en temps réel dans un formulaire de contenu.
  • Une application qui doit respecter la vie privée : les images de santé ou les documents financiers ne quittent pas le poste du client.

WebGPU (via executionProviders: ["webgpu"]) permet d'exploiter le GPU du client ; les gains sur un ResNet18 sont d'un facteur 5 à 10 sur un MacBook récent. Cette voie est jeune mais avance vite ; elle mérite d'être connue même si votre déploiement principal reste côté serveur.

Ce que ce cours n'a pas traité

Trois sujets délibérément laissés hors périmètre, avec le renvoi vers le module qui les traite.

  • La distribution du modèle (registre, versionnage, canari) est un sujet MLOps couvert au cours 33.
  • La surveillance en production — dérive de distribution, taux d'erreurs, alertes — est traitée au cours 32.
  • L'automatisation du pipeline export → optimisation → quantification → mesure est le sujet du cours 34.

ONNX Runtime n'est qu'un maillon dans cette chaîne, mais c'est celui qui rend possible tout le reste : sans format d'échange, aucune de ces pratiques MLOps ne tient.

En résumé

  • Une session ONNX Runtime unique, partagée entre requêtes, est le bon design : la créer par requête gaspille des centaines de millisecondes et de la mémoire.
  • Le prétraitement doit accompagner le modèle — normalisation, ordre des canaux, recadrage — versionné à côté du fichier .onnx pour survivre aux changements d'équipe.
  • Le regroupement en lot côté service tire parti du matériel sur GPU ; en dessous de 5 à 10 millisecondes de latence acceptable, la mise en lot est indispensable.
  • ONNX Runtime Web exécute le même fichier .onnx côté navigateur via WebAssembly ou WebGPU ; c'est la voie pour la vie privée et pour les applications qui ne tolèrent pas la latence réseau.

Le récapitulatif final relie les dix modules, propose une liste de contrôle d'un export sain, et annonce l'examen de 40 questions et l'attestation.