Le portage du modèle Gemma-4 31B de Google sur un AWS Inferentia2 inf2.24xlarge a produit une correspondance parfaite, jeton par jeton, par rapport à la référence CPU — et pourtant, chaque phrase générée était un charabia. L'écart entre « correspondre » et « fonctionner » sert désormais d'avertissement pour quiconque tente de faire tenir des LLM massifs sur les puces d'inférence personnalisées d'Amazon.
Pourquoi une correspondance jeton par jeton ne suffit pas
Le développeur a comparé chaque jeton de sortie de l'appareil Inferentia avec le jeton produit par une exécution du modèle sur CPU. Les flux étaient identiques, de sorte que le matériel semblait avoir reproduit exactement l'implémentation de référence. En réalité, les deux flux injectaient un prompt malformé dans un modèle privé de son chat template et doté de mauvais marqueurs de tour de parole. L'absence de template a plongé le modèle dans une boucle infinie, recrachant du non-sens. Le matériel a fait son travail : il a reproduit un bug qui existait dans le code de référence.
La leçon est simple : SEQ_MATCH (égalité séquentielle des jetons) ne signifie pas exactitude. Si l'implémentation de référence est défectueuse, une réplique matérielle fidèle hérite du même échec. La validation doit aller au-delà de la parité au niveau des jetons ; elle nécessite des tests fonctionnels de bout en bout avec des entrées correctement formatées.
Des buffers se faisant passer pour des paramètres
Lors de la phase de chargement, le chargeur de modèle a ignoré un composant appelé layer_scalar. Le code enregistrait cet objet en tant que buffer plutôt qu'en tant que paramètre dans la définition du modèle PyTorch. Les buffers sont des tenseurs statiques que l'entraînement ne met pas à jour, et de nombreux chargeurs les ignorent lors de la conversion vers des formats compatibles Neuron. Son omission a laissé les facteurs d'échelle de plusieurs couches à leurs valeurs par défaut, faussant les calculs sur l'ensemble du réseau. Aucune erreur n'a été signalée ; le modèle a été compilé et le pipeline d'inférence s'est exécuté, mais les résultats numériques étaient erronés.
Pour toute personne transférant de grands modèles sur Inferentia, auditez chaque tenseur non-paramètre. Même si un tenseur n'est pas destiné à être appris, il peut rester essentiel pour le calcul correct de la passe avant (forward-pass). Vérifier manuellement l'inclusion des buffers peut prévenir des erreurs d'échelle silencieuses qui sont autrement difficiles à diagnostiquer.
Volatilité des instances Spot et compilation de 39 minutes
Faire tourner un modèle de 31 milliards de paramètres sur une instance Spot semble économique, mais les économies s'accompagnent d'événements de récupération imprévisibles. Le temps de compilation du développeur — environ 39 minutes pour traduire le modèle en code compatible Neuron — s'est volatilisé lorsque l'AWS a récupéré l'instance. Pour survivre aux interruptions, il a mis en place un filet de sécurité à trois volets :
- ModelBuilder maintenait l'utilisation de la mémoire dans la limite de 384 Go de l'hôte, évitant ainsi les plantages qui imposeraient un redémarrage.
- Un mirroring S3 immédiat des fichiers de poids bruts et des « neffs » compilés (fichiers exécutables Neuron) permettait à une nouvelle instance de reprendre exactement là où la précédente s'était arrêtée.
- Un polleur multi-région scannait les régions AWS à la recherche de capacités Spot disponibles et lançait une nouvelle instance dès qu'une capacité apparaissait.
Ces étapes ont transformé une compilation fragile à point de défaillance unique en un pipeline résilient capable de survivre à l'instabilité des marchés Spot.
Pièges du sharding avec des configurations d'attention mixtes
Gemma-4 31B utilise deux configurations d'attention. Certaines couches emploient quatre têtes clé-valeur (KV), d'autres un nombre différent. Diviser le modèle uniformément sur huit rangs (ranks) parallèles échoue lorsqu'un nombre de têtes KV d'une couche ne se divise pas proprement. Tenter de partitionner (shard) une couche à 4 têtes sur huit rangs forcerait chaque rang à gérer une demi-tête — une impossibilité mathématique qui déclenche des erreurs de dimension (shape mismatches) et des erreurs d'exécution.
La solution a consisté à répliquer les couches partitionnées globalement (celles ayant des nombres de têtes compatibles) sur tous les rangs, et à ne partitionner que les couches « glissantes » dont le nombre de têtes permettait une division égale. Cette stratégie hybride a permis de conserver l'efficacité du parallélisme de tenseurs tout en évitant la division illégale des têtes KV, éliminant ainsi les erreurs de parallélisation de tenseurs qui avaient entravé les tentatives précédentes.
À retenir
Porter un LLM géant sur Inferentia est bien plus qu'un simple exercice de compilation et d'exécution. Cela exige des tests fonctionnels rigoureux au-delà de l'égalité des jetons, une vérification méticuleuse que chaque tenseur — paramètre ou buffer — est correctement géré, et une stratégie de déploiement qui anticipe la récupération des instances Spot. Enfin, le partitionnement (sharding) doit respecter la géométrie d'attention interne du modèle ; sinon, le parallélisme qui promet de la vitesse devient une source d'échec silencieux.
