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.
- 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.
- Performance : les moteurs (LibTorch, ONNX Runtime) appliquent des optimisations spécifiques à l'inférence : fusion d'opérations, quantification, choix d'algorithmes.
- 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 pastrace 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
.pydu modèle :torch.jit.loadlit un artefact autonome.
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
| Aspect | TorchScript | ONNX |
|---|---|---|
| Portabilité | LibTorch (C++, Java) | Multi-moteurs, multi-langages |
| Fidélité | Élevée avec script, correcte avec trace | Bonne, dépend des opérateurs supportés |
| Optimisation | Modérée | Excellente avec ONNX Runtime, TensorRT |
| Complexité | Faible | Modérée (opset, axes dynamiques) |
| Cas d'usage | Service PyTorch pur, embarqué mobile PyTorch | Multi-cadre, edge, mobile Apple/Android |
En résumé
tracecapture l'exécution observée, pas la logique ;scriptanalyse 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_axespour 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.