Aller au contenu principal

Module 10 — TorchScript, ONNX et mise en service

Un modèle entraîné n'a de valeur que s'il est utilisable en production. Rester en Python impose un interpréteur, un environnement conda ou pip reproduit à l'octet près, une dépendance à torch complet. Ce module présente les deux formats standards pour exporter un modèle PyTorch et le servir sans son code : TorchScript (natif PyTorch) et ONNX (interopérable). Le fil rouge est celui du module précédent : le ResNet18 affiné sur Fashion-MNIST.

Pourquoi exporter au lieu de servir Python

Trois raisons cumulatives.

  1. Portabilité : un artefact TorchScript ou ONNX se charge dans un moteur d'inférence C++, Java, JavaScript, Go, sans jamais démarrer d'interpréteur Python.
  2. Performance : les moteurs (LibTorch, ONNX Runtime) appliquent des optimisations spécifiques à l'inférence : fusion d'opérations, quantification, choix d'algorithmes.
  3. Reproductibilité : l'artefact contient le graphe du modèle. Un successeur qui reprend le projet six mois plus tard n'a pas besoin du code Python pour l'exécuter.

torch.jit.trace : deux minutes, un piège

trace exécute le modèle avec une entrée exemple et enregistre les opérations effectivement exécutées. Le graphe résultant reproduit exactement ce qui s'est passé sur cette entrée.

import torch

modele.eval()
exemple = torch.randn(1, 3, 224, 224)
modele_trace = torch.jit.trace(modele, exemple)
modele_trace.save("resnet18_fashion.pt")

Simple, silencieux, rapide. Le piège est majeur : trace capture uniquement le chemin d'exécution observé. Un if x.max() > 0 dans le forward est réduit à la branche prise sur l'exemple donné. Une taille de lot variable n'est pas capturée si le graphe dépend d'elle.

trace capture, ne raisonne pas

trace produit un TracerWarning quand il détecte un contrôle de flux suspect. Toujours le lire : c'est la seule protection contre un modèle exporté qui ne fait pas ce que l'original faisait. En cas de doute, préférer script.

torch.jit.script : plus lent à écrire, plus fidèle

script analyse le code source du forward en une variante statiquement typée de Python. Le graphe résultant contient les branchements, les boucles, les types. C'est plus contraignant — pas de *args, types annotés — mais fidèle.

class ClassificateurScriptable(nn.Module):
def __init__(self):
super().__init__()
self.backbone = resnet18(weights=None)
self.backbone.fc = nn.Linear(512, 10)

def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.backbone(x)

modele_scripte = torch.jit.script(ClassificateurScriptable())
modele_scripte.save("modele_scripte.pt")

Règle pratique : trace pour un modèle purement séquentiel sans branchement (ResNet, VGG, la plupart des CNN). script dès qu'il y a du contrôle de flux dans forward, ou pour un modèle destiné à un usage large (RNN, tokenisateur intégré, prétraitement conditionnel).

Le contrôle numérique en trois lignes

Un modèle exporté est présumé cassé jusqu'à preuve du contraire. Comparer la sortie de l'original et de l'export sur un lot fixé est rapide et détecte immédiatement une régression.

modele.eval()
modele_trace.eval()

entree = torch.randn(4, 3, 224, 224)
with torch.no_grad():
sortie_origine = modele(entree)
sortie_export = modele_trace(entree)

max_ecart = (sortie_origine - sortie_export).abs().max().item()
print(f"écart max : {max_ecart:.2e}")
assert torch.allclose(sortie_origine, sortie_export, atol=1e-5)

Un écart supérieur à 1e-4 sur des sorties float32 signale un problème : opération non capturée, mode train oublié, différence de BatchNorm. C'est la deuxième fonction à écrire, juste après l'export.

Charger sans le code Python : torch.jit.load

Le fichier .pt d'un export TorchScript est autonome. Sur une machine où l'on installe seulement torch, sans le fichier .py du modèle, le chargement fonctionne.

import torch

modele = torch.jit.load("resnet18_fashion.pt")
modele.eval()

entree = torch.randn(1, 3, 224, 224)
with torch.no_grad():
logits = modele(entree)
predictions = logits.argmax(dim=1)

C'est déjà un scénario de production minimaliste : ce script tourne en service, écoute des requêtes, et n'a jamais besoin du dépôt d'entraînement.

ONNX : sortir de l'écosystème PyTorch

ONNX (Open Neural Network Exchange) est un format ouvert compris par de nombreux moteurs : ONNX Runtime (Microsoft), TensorRT (NVIDIA), Core ML (Apple), plusieurs cadres mobiles et embarqués. C'est le passage obligé quand on doit servir en dehors de LibTorch.

import torch

modele.eval()
exemple = torch.randn(1, 3, 224, 224)

torch.onnx.export(
modele,
exemple,
"resnet18_fashion.onnx",
input_names=["entree"],
output_names=["logits"],
dynamic_axes={
"entree": {0: "batch"}, # taille de lot variable
"logits": {0: "batch"},
},
opset_version=17,
)

Deux paramètres à comprendre. dynamic_axes déclare quelles dimensions sont variables ; sans lui, le graphe fige la taille de lot à celle de l'exemple, ce qui rend l'inférence par lots impossible. opset_version choisit la version du jeu d'opérateurs ; 17 est un compromis raisonnable en 2026 entre compatibilité et couverture.

Vérifier l'export ONNX

Comme pour TorchScript, on vérifie numériquement. onnxruntime exécute le modèle exporté en dehors de PyTorch et on compare.

import numpy as np
import onnxruntime as ort

session = ort.InferenceSession("resnet18_fashion.onnx", providers=["CPUExecutionProvider"])
entree_np = np.random.randn(4, 3, 224, 224).astype(np.float32)

sortie_ort = session.run(["logits"], {"entree": entree_np})[0]
sortie_torch = modele(torch.from_numpy(entree_np)).detach().numpy()

print("écart max :", np.abs(sortie_torch - sortie_ort).max())

Un micro-service d'inférence en 30 lignes

Pour donner une idée concrète — le module 37 du parcours reprend ce sujet en profondeur. Ici, on charge le modèle TorchScript et on répond en HTTP.

# service.py
from fastapi import FastAPI, UploadFile
from io import BytesIO
from PIL import Image
import torch
from torchvision import transforms

app = FastAPI()
modele = torch.jit.load("resnet18_fashion.pt")
modele.eval()

transform = transforms.Compose([
transforms.Grayscale(num_output_channels=3),
transforms.Resize(224),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
classes = ["T-shirt", "Pantalon", "Pull", "Robe", "Manteau",
"Sandale", "Chemise", "Basket", "Sac", "Bottine"]

@app.post("/predire")
async def predire(fichier: UploadFile):
image = Image.open(BytesIO(await fichier.read())).convert("L")
tenseur = transform(image).unsqueeze(0)
with torch.no_grad():
logits = modele(tenseur)
i = int(logits.argmax(dim=1).item())
return {"classe": classes[i], "index": i}

Trois points essentiels de ce service.

  • La transformation d'entrée est dans le service, identique à celle du module 9. Si elle diffère, le modèle rend n'importe quoi.
  • torch.no_grad() est là aussi, pour économiser mémoire et temps à chaque requête.
  • Aucune dépendance au fichier .py du modèle : torch.jit.load lit un artefact autonome.
Le prétraitement suit toujours le modèle

La leçon centrale de ce module : ce qui n'est pas dans l'artefact n'existe pas. Une normalisation faite dans le script d'entraînement mais oubliée dans le service produit un modèle « qui ne fonctionne pas en production » alors que rien n'est cassé dans PyTorch. Documenter la transformation à côté du fichier .pt ou .onnx, ou l'inclure dans un module scripté enveloppe qui contient le prétraitement, est la seule protection.

Comparaison rapide

AspectTorchScriptONNX
PortabilitéLibTorch (C++, Java)Multi-moteurs, multi-langages
FidélitéÉlevée avec script, correcte avec traceBonne, dépend des opérateurs supportés
OptimisationModéréeExcellente avec ONNX Runtime, TensorRT
ComplexitéFaibleModérée (opset, axes dynamiques)
Cas d'usageService PyTorch pur, embarqué mobile PyTorchMulti-cadre, edge, mobile Apple/Android

En résumé

  • trace capture l'exécution observée, pas la logique ; script analyse le code source et préserve le contrôle de flux.
  • Contrôler numériquement l'export avant toute mise en service : torch.allclose(sortie_origine, sortie_export, atol=1e-5).
  • ONNX ouvre la porte à des moteurs d'inférence externes ; ne pas oublier dynamic_axes pour une taille de lot variable.
  • Le prétraitement doit accompagner l'artefact, pas rester dans le script d'entraînement ; sinon le service rend des prédictions aberrantes sans qu'aucune erreur ne soit levée.

Le récapitulatif final relie les dix modules, souligne les fils qui les traversent et prépare l'examen de 40 questions et l'attestation.