Titre : Le bug résidait dans un seul axe
Un bug caché dans l'implémentation du chunked-scan de PyTorch permet à certains modèles hybrides d'apercevoir les tokens futurs lorsque les noyaux fusionnés rapides (fast fused kernels) sont absents, compromettant des checkpoints tels que Zamba2-1.2B et Nemotron-H-8B. La faille provient d'une seule réduction effectuée selon le mauvais axe, et elle se manifeste sur le chemin d'exécution sur lequel la plupart des pipelines CI et des exécutions CPU s'appuient.
Pourquoi ce bug est important
Les modèles hybrides et les modèles d'espace d'état (state-space models) ne dépendent plus uniquement de l'attention ; ils utilisent également des récurrences linéaires, des scans et des convolutions. Une opération de masquage qui bloque les tokens futurs dans la matrice d'attention ne garantit pas automatiquement la causalité pour ces opérations supplémentaires. Lorsque les noyaux fusionnés rapides qui traiteraient normalement le scan correctement sont manquants, PyTorch se rabat sur un chemin de chunked-scan en pur Python. Ce mode de secours contient une confusion d'axes (axis-mixup), permettant aux informations des positions ultérieures de circuler vers l'arrière pendant l'inférence.
La fuite est silencieuse. Elle ne fait pas planter le modèle et ne génère pas d'erreur évidente. Au lieu de cela, elle rend la perte (loss) et la perplexité artificiellement basses car le modèle triche concrètement — en voyant les tokens mêmes qu'il est censé prédire. Toute évaluation en aval qui se fie à ces métriques est donc bâtie sur des fondations fragiles.
Comment le problème a été découvert
Des chercheurs ont comparé deux passes avant (forward passes) à travers le même modèle :
- Une séquence de tokens aléatoires.
- La séquence identique avec un seul token modifié.
Ils ont mesuré la différence dans les états cachés (hidden states) couche par couche. L'inspection du masque n'a rien signalé, mais un audit par couche ayant injecté des fautes a détecté 192 erreurs sur 192 injectées, confirmant la fuite.
Un examen rapide de la bibliothèque transformers a révélé :
- Zamba2-1.2B présente une fuite lorsque la taille du chunk est fixée à 256.
- Nemotron-H-8B présente une fuite avec une taille de chunk de 128.
- La plupart des autres checkpoints ne montraient aucune fuite dans les mêmes conditions.
Le défaut réside dans le chemin de code du chunked-scan qui s'exécute chaque fois que les noyaux fusionnés optionnels sont absents. Cela inclut :
- Toutes les exécutions sur CPU.
- Les environnements GPU sans les packages de noyaux fusionnés spécifiques.
- Les installations PyTorch standard qui omettent les dépendances supplémentaires.
Comme le bug n'apparaît que lorsque ces noyaux sont absents, il peut survenir dans les environnements de CI et sur les CPU.
Qui est à risque et quel est le coût
Toute équipe qui entraîne, affine (fine-tune) ou évalue des modèles hybrides sans les noyaux fusionnés risque de publier des chiffres de performance gonflés. L'amélioration apparente de la perte ou de la perplexité est illusoire ; le modèle a concrètement « regardé en avant ». Pour les groupes de recherche, cela peut conduire à des affirmations trompeuses sur des résultats de pointe (state-of-the-art). Pour les déploiements commerciaux, cela peut causer des erreurs en aval dans les tâches de génération qui n'ont jamais été réellement apprises.
L'argument opposé
Certains développeurs soutiennent qu'un masque causal correctement appliqué est suffisant pour empêcher toute fuite de tokens futurs. Le bug infirme cette idée : les scans, les convolutions et certaines couches de normalisation peuvent contourner entièrement le masque. La confusion d'axes dans le chunked-scan montre que la causalité doit être appliquée partout où les données circulent, et pas seulement dans la matrice d'attention.
Ce qu'il faut surveiller ensuite
- Hygiène des dépendances : Installez les packages de noyaux fusionnés rapides sur tous les nœuds d'entraînement et d'inférence, en particulier dans les pipelines CI.
- Scripts d'audit : Suivez l'audit en deux étapes recommandé par les découvreurs :
- Injectez une faute connue dans le checkpoint que vous testez comme contrôle positif.
- Exécutez des séquences plus longues que la taille du chunk ou de la fenêtre du modèle ; les séquences courtes ne révéleront jamais la fuite.
L'exécution de l'audit sur n'importe quel checkpoint hybride ou d'espace d'état avec une longueur de séquence dépassant la taille du chunk révélera si le modèle est toujours vulnérable.
À retenir : Même un modèle qui réussit tous les tests standards peut tricher silencieusement lorsque le chemin d'exécution se rabat sur une implémentation buggée. Vérifier la présence des noyaux fusionnés rapides — et tester explicitement la fuite au-delà des masques d'attention — sont désormais des étapes essentielles avant de faire confiance aux métriques de n'importe quel modèle hybride.
