Aller au contenu principal

Module 5 — Optimisation du graphe

Un modèle qui vient de sortir de torch.onnx.export ou de tf2onnx est fidèle, pas optimisé. Il contient souvent des Transpose redondants, des multiplications par un scalaire séparées d'une convolution, des embranchements de Cast qui coûtent sans rien apporter. ONNX Runtime propose plusieurs niveaux d'optimisation de graphe qui, appliqués une fois, accélèrent l'inférence de 20 % à 60 % sans changer la précision numérique. Ce module explique ce qu'on gagne, ce qu'on risque et comment mesurer.

Trois familles de transformations

Les optimisations de graphe se rangent en trois familles.

Le pliage de constantes (constant folding) élimine tout sous-graphe dont les entrées sont toutes des initialiseurs. Une expression Add(Mul(constante_A, constante_B), constante_C) devient un seul tenseur constant précalculé. C'est le cas d'école quand des couches de normalisation avec moyenne et écart-type fixes suivent une convolution : le graphe garde une seule multiplication et une seule addition.

La fusion d'opérateurs remplace un motif reconnu par un opérateur composite plus efficace. Les motifs les plus courants :

  • Conv + BN + ReLU fusionné en un FusedConv qui applique tout en un seul kernel.
  • MatMul + Add fusionné en Gemm (matrice générale avec biais).
  • LayerNorm remonté en un seul opérateur au lieu d'une chaîne Reduce, Sub, Mul, Add.
  • Attention fusionné pour les transformeurs en opset récent.

L'élimination de redondances supprime les nœuds inutiles : Identity intermédiaire, Transpose(Transpose(x)) = x, Cast qui convertit vers le même type. Ces optimisations sont locales et ne changent jamais la sémantique du graphe.

Les niveaux d'ONNX Runtime

ONNX Runtime expose ces optimisations sous quatre niveaux, appliqués séquentiellement.

NiveauNomContenu
0ORT_DISABLE_ALLAucune optimisation ; utile pour le débogage
1ORT_ENABLE_BASICPliage de constantes, élimination des Identity, fusion des Add/Mul triviaux
2ORT_ENABLE_EXTENDEDFusions plus lourdes : FusedConv, Gemm, LayerNorm, Attention
3ORT_ENABLE_ALLToutes les optimisations, y compris celles spécifiques au fournisseur d'exécution

Le paramètre est fixé à la création de la session :

import onnxruntime as ort

options = ort.SessionOptions()
options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL

session = ort.InferenceSession(
"resnet18_fashion.onnx",
sess_options=options,
providers=["CPUExecutionProvider"],
)

ORT_ENABLE_ALL est le défaut recommandé en production. Le niveau 2 sert quand on veut isoler un bug qu'on soupçonne dans les optimisations spécifiques au fournisseur. Le niveau 0 sert de référence quand on veut mesurer l'apport des optimisations.

Mesurer le gain sur le ResNet18

Un protocole minimal, sur le fil rouge : mesurer la latence d'inférence sur un lot de 32 images avec chacun des quatre niveaux.

import time
import numpy as np
import onnxruntime as ort

def latence_moyenne(niveau, n_iter=200, echauffement=20):
options = ort.SessionOptions()
options.graph_optimization_level = niveau
session = ort.InferenceSession(
"resnet18_fashion.onnx",
sess_options=options,
providers=["CPUExecutionProvider"],
)
entree = np.random.randn(32, 3, 224, 224).astype(np.float32)

for _ in range(echauffement):
session.run(None, {"entree": entree})

debut = time.perf_counter()
for _ in range(n_iter):
session.run(None, {"entree": entree})
fin = time.perf_counter()
return (fin - debut) / n_iter * 1000 # en millisecondes

for niveau in [ort.GraphOptimizationLevel.ORT_DISABLE_ALL,
ort.GraphOptimizationLevel.ORT_ENABLE_BASIC,
ort.GraphOptimizationLevel.ORT_ENABLE_EXTENDED,
ort.GraphOptimizationLevel.ORT_ENABLE_ALL]:
ms = latence_moyenne(niveau)
print(f"{niveau.name:22s} : {ms:.2f} ms/lot")

Sur un CPU x86 récent, l'ordre attendu sur ce modèle est proche de :

  • ORT_DISABLE_ALL : 62 ms/lot
  • ORT_ENABLE_BASIC : 51 ms/lot
  • ORT_ENABLE_EXTENDED : 39 ms/lot
  • ORT_ENABLE_ALL : 37 ms/lot

Le saut principal se fait au niveau 2, quand FusedConv remplace le triplet Conv + BN + ReLU. Le niveau 3 apporte peu sur CPU pur, davantage quand des fournisseurs comme CUDA ou TensorRT sont dans la boucle.

Sauvegarder le graphe optimisé

L'optimisation prend quelques centaines de millisecondes au démarrage de la session. Pour un service qui recharge fréquemment le modèle, ou pour figer le graphe une fois pour toutes, on demande à ONNX Runtime de sauvegarder le graphe optimisé sur disque.

options = ort.SessionOptions()
options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_EXTENDED
options.optimized_model_filepath = "resnet18_fashion_optimized.onnx"

# Créer la session déclenche l'optimisation et écrit le fichier
_ = ort.InferenceSession(
"resnet18_fashion.onnx",
sess_options=options,
providers=["CPUExecutionProvider"],
)

Attention à un détail important : le graphe sauvegardé peut contenir des fusions spécifiques à un fournisseur. Un fichier sauvegardé avec le fournisseur CUDA ne se rechargera pas correctement sur une machine CPU. Pour un fichier portable, on limite le niveau à ORT_ENABLE_EXTENDED — les fusions du niveau 3 spécifiques au fournisseur ne sont pas incluses.

La portabilité du graphe optimisé

Un graphe issu de ORT_ENABLE_ALL avec CUDA activé peut faire référence à des opérateurs CUDA propriétaires (FusedConvBiasAct par exemple). Servi sur une machine CPU seule, il refusera de se charger. Pour un artefact portable, on optimise avec le fournisseur cible réel et on documente la contrainte. Pour un artefact passe-partout, on reste à ORT_ENABLE_EXTENDED.

Le pliage de constantes vu par Netron

Après optimisation, ouvrir le fichier optimisé dans Netron révèle visuellement ce qui a changé. Sur le ResNet18 :

  • Les blocs Conv - BN sont remplacés par un Conv unique qui absorbe les paramètres de la normalisation dans son biais et son échelle.
  • Les branches courtes (raccourcis résiduels) sont conservées telles quelles ; on ne fusionne pas ce qui traverse plusieurs chemins.
  • La couche finale Linear + Softmax reste en deux nœuds : le pliage ne fusionne que les motifs qu'il reconnaît.

Cette relecture visuelle est utile parce qu'elle permet de repérer ce qui n'a pas été fusionné. Un Conv isolé, suivi d'un BN mais séparé par un Reshape intermédiaire, échappe à la fusion. Un ajustement du code source du modèle peut alors débloquer la fusion et gagner quelques millisecondes.

Les limites : ce que l'optimisation ne fait pas

L'optimisation de graphe est conservative. Elle ne change ni la précision numérique, ni la structure du modèle : elle ne quantifie pas, ne fusionne pas d'opérateurs qui changeraient la sortie, ne coupe pas de branches. Trois choses qu'elle ne fait donc pas :

  • Convertir le modèle en INT8 (module 6 : quantification).
  • Compiler pour un accélérateur spécifique (module 7 : fournisseurs).
  • Éliminer un Softmax que le service applique lui-même en aval — c'est une optimisation métier qui reste manuelle.

Une optimisation de graphe qui gagne 40 % sur CPU est un excellent point de départ, mais il faut ensuite décider si l'on va plus loin en acceptant une perte de précision (quantification) ou une dépendance matérielle (TensorRT).

Vérifier que l'optimisation n'a rien cassé

L'optimisation est réputée équivalente numériquement. Le mot « réputée » signifie que la vérification du module 4 se refait après optimisation : atol=1e-4 doit tenir.

sortie_avant = session_niveau_0.run(None, {"entree": entree_np})[0]
sortie_apres = session_niveau_3.run(None, {"entree": entree_np})[0]

ecart = np.abs(sortie_avant - sortie_apres).max()
assert ecart < 1e-4, f"L'optimisation a introduit un écart de {ecart:.2e}"

Un écart de 1e-6 à 1e-5 est normal : les fusions changent l'ordre des additions. Un écart plus grand signale un bug de fusion, qu'il faut remonter au projet ONNX Runtime.

En résumé

  • Trois familles d'optimisations : pliage de constantes, fusion d'opérateurs et élimination de redondances ; toutes sont mathématiquement équivalentes au graphe d'origine.
  • ONNX Runtime expose quatre niveaux ; ORT_ENABLE_ALL est le défaut en production, mais on reste à ORT_ENABLE_EXTENDED pour un artefact portable entre fournisseurs.
  • Le graphe optimisé se sauvegarde avec optimized_model_filepath ; il faut alors documenter le fournisseur cible car des fusions spécifiques peuvent bloquer le rechargement ailleurs.
  • La vérification numérique du module 4 se rejoue après optimisation ; un atol supérieur à 1e-4 sur du float32 signale une régression à investiguer.

Le module suivant descend d'un cran de précision : on quantifie le modèle en INT8, dynamiquement puis avec calibration, et on mesure ce qu'on gagne et ce qu'on perd.