Loading
L'int8 naïf plafonnait la parité d'attention à 2,97e-2 — trente fois pire que la barre. Où dépenser la précision : voilà toute l'ingénierie.
Ouvrir l'instrumentLa façon évidente d'exécuter un transformeur dans un navigateur, c'est ONNX Runtime Web. C'est ce que ce projet utilisait d'abord, et cela fonctionnait.
Cela coûtait aussi un runtime WebAssembly dans le bundle, une exception de Content-Security-Policy pour autoriser la compilation WASM, et une dépendance dont je ne maîtrisais pas le cycle de publication — le tout pour exécuter quatre blocs de transformeur sur une phrase d'au plus vingt-quatre tokens.
Je l'ai donc supprimé et j'ai écrit la passe avant à la main.
L'idée décisive : un modèle à quatre couches n'a pas besoin d'un moteur d'inférence généraliste. Il lui faut environ sept cents lignes d'arithmétique sur
Float32Array, et toutes les contraintes imposées par le moteur disparaissent.
Tout le runtime est désormais du TypeScript simple : embeddings, layer norm, trois projections linéaires par couche, produits scalaires mis à l'échelle, un softmax par ligne, la somme pondérée des values, un feed-forward à deux couches avec GELU exact, et les connexions résiduelles. Environ 770 lignes, tokenizer WordPiece compris.
Pour une phrase de douze mots, cela prend des dizaines de millisecondes. Le tout tourne dans un Web Worker de type module : le thread principal n'est jamais bloqué et l'animation reste fluide pendant l'exécution.
Ce que cela a rapporté :
'wasm-unsafe-eval'. La CSP s'est resserrée à script-src 'self', et le test navigateur vérifie désormais que l'en-tête ne contient pas l'exception qu'il exigeait auparavant.Écrire la passe avant a pris un après-midi. La faire coïncider avec PyTorch a pris nettement plus longtemps, et c'est la seule partie qui mérite un article.
Les poids sont livrés quantifiés — int8, avec une échelle par canal de sortie plutôt qu'une par tenseur, car une échelle unique perd trop sur les couches feed-forward. Les échelles par canal coûtent quatre octets par ligne et les valent toutes.
La barre : erreur d'attention au pire cas inférieure à 1e-3 face à une référence PyTorch fp32, ce qui rend fiable la troisième décimale affichée. Si l'interface écrit 0,612, la réponse du modèle doit arrondir à 0,612.
L'int8 intégral naïf mesurait 2,97e-2. Trente fois la barre. L'attention était visiblement fausse.
Erreur d'attention au pire cas face à PyTorch fp32, mesurée à chaque étape de la conception de la quantification. La barre est à 1e-3. Chiffres issus de engine/bert.parity.test.ts.
Ce qui a corrigé cela n'est pas une précision uniforme, mais le fait de découvrir où l'erreur devient visible et de n'y dépenser des octets que là.
Embeddings de position et de type en fp32. Ce sont de minuscules tables, et leur erreur entre directement dans les logits de la couche 0, avant qu'aucun bloc n'ait tourné. En int8, elles maintenaient à elles seules la parité à 9,4e-3 : le réseau n'avait encore rien fait et était déjà grossier.
Q et K en int16. Leur erreur est amplifiée par l'exponentielle du softmax. Une petite perturbation d'un logit devient une grande perturbation d'une probabilité.
V et la sortie d'attention en int16. Elles écrivent dans le flux résiduel : leur erreur se propage à toutes les couches suivantes.
Feed-forward en int16 pour toutes les couches sauf la dernière. Le feed-forward final n'alimente aucune attention — rien en aval ne le lit — il reste donc en int8, gratuitement.
Mesure finale : 1,218e-4, confortablement sous la barre, pour un chargement à froid de 6,32 Mio.
Chaque barreau ci-dessus est une hypothèse mesurée, pas une règle empirique. « Tout quantifier en int8 » et « tout garder en fp32 » sont deux réponses faciles et fausses : la première casse la sortie, la seconde gaspille des mégaoctets sur des tables qui s'en moquent.
La bonne question n'est jamais « quelle précision ? », mais « de quels tenseurs l'erreur s'échappe-t-elle, et où est-elle amplifiée ? ». Le softmax amplifie. Les résiduelles propagent. Un feed-forward en cul-de-sac ne fait ni l'un ni l'autre.
Cela ne se sait qu'en mesurant couche par couche. Le test de parité rapporte l'erreur par couche pour exactement cette raison : si la couche 0 est déjà grossière, le chemin des embeddings est en cause ; si l'erreur croît avec la profondeur, l'arithmétique des blocs dérive. Deux bugs différents, deux correctifs différents — et un seul chiffre agrégé les cacherait tous les deux.
Une parité à 1,2e-4 dit que cette implémentation s'accorde avec PyTorch sur les phrases testées. Elle ne dit pas que l'arithmétique est correcte en général : ce serait une affirmation sur toutes les entrées, et aucun test fini ne l'établit.
D'où le fait que la suite de tests ne s'arrête pas à la phrase de référence, et qu'une seconde implémentation existe, sans aucun code commun avec celle-ci. La partie 8 en parle — et du bug qu'elle a trouvé.
Ensuite : télécharger vingt lignes d'une table de trente et un mégaoctets.
Plus à lire
Choisir la tête la plus confiante sélectionnait une tête positionnelle à chaque fois. Elle rapportait « it → was, 0,962 » — la réponse qu'elle donne pour chaque mot de chaque phrase.
La table d'embeddings pèse 31 254 528 octets. Une phrase de douze mots en utilise environ 20 Ko. HTTP Range fait que c'est la seule partie qui circule.
Trois matrices, un produit scalaire, un softmax. La ligne somme toujours à 1 — une contrainte que le modèle doit satisfaire, pas un résultat qu'il rapporte.