RADLADS: Rapid Attention Distillation to Linear Attention Decoders at Scale
Auteurs originaux : Daniel Goldstein, Eric Alcaide, Janna Lu, Eugene Cheah
Auteurs originaux : Daniel Goldstein, Eric Alcaide, Janna Lu, Eugene Cheah
Article original sous licence CC BY 4.0 (http://creativecommons.org/licenses/by/4.0/). ✨ Ceci est une explication générée par l'IA de l'article ci-dessous. Elle n'a pas été rédigée ni approuvée par les auteurs. Pour une précision technique, consultez l'article original. Lire la clause de non-responsabilité complète
Résumé Technique : RADLADS
Énoncé du Problème
Bien que les variantes de l'attention linéaire (telles que RWKV et Mamba) offrent des avantages significatifs par rapport aux transformeurs d'attention softmax traditionnels — spécifiquement un temps d'inférence de O(1) par jeton et une utilisation constante de la mémoire en évitant les caches Clé-Valeur (KV cache) — l'entraînement de modèles linéaires à grande échelle à partir de zéro reste prohibitif en termes de coûts. Les modèles de transformeurs de pointe (SoTA) nécessitent souvent un entraînement sur plus de dix mille milliards de jetons, un coût inaccessible à la plupart des organisations. Les méthodes existantes pour convertir les transformeurs softmax pré-entraînés en modèles linéaires ou récurrents (par exemple, T2R, SUPRA, LoLCats, MOHAWK) ont historiquement nécessité des nombres massifs de jetons (allant de 20 Md à 100 Md+ de jetons) ou ont abouti à de faibles performances en aval, particulièrement sur des benchmarks comme MMLU. De plus, de nombreuses tentatives de conversion reposent sur des architectures hybrides (conservant une partie de l'attention softmax) ou échouent à égaler la qualité des modèles enseignants originaux.
Méthodologie
Les auteurs présentent RADLADS (Rapid Attention Distillation to Linear Attention Decoders at Scale), un protocole conçu pour convertir les transformeurs d'attention softmax en modèles de décodeurs à attention linéaire avec un coût computationnel minimal. Le processus comprend trois étapes principales et des adaptations architecturales spécifiques :
1. Le Protocole RADLADS
La conversion est un processus en trois étapes ne nécessitant que 350 à 700 millions de jetons (moins de 0,005 % des données de pré-entraînement de l'enseignant) :
- Configuration (Transfert des poids d'attention) : Les poids liés à l'attention (Wq,Wk,Wv,Wo) du modèle enseignant sont transférés directement vers les couches de mélange de séquence du modèle étudiant. Les autres poids sont initialisés via des méthodes de pré-entraînement standard ou configurés pour imiter le comportement de l'enseignant sans effet immédiat.
- Étape 1 : Alignement de l'état caché de l'attention : Les couches de mélange de séquence du modèle étudiant sont entraînées pour approximer les sorties d'état caché des couches d'attention correspondantes de l'enseignant. Cela est réalisé à l'aide d'un objectif de distance L2 (ou MSE). Les auteurs ont constaté que l'utilisation d'un noyau d'attention linéaire à porte (Gated Linear Attention), en supprimant les termes de décroissance "off-by-one" et de bonus, permet un ajustement plus proche des états cachés de l'enseignant que le RWKV-6 standard.
- Étape 2 : Distillation de la connaissance : L'ensemble du modèle étudiant est entraîné pour approximer les logits de sortie de l'enseignant en utilisant une perte de divergence de Kullback-Leibler (KL). Les auteurs émettent l'hypothèse que la connaissance factuelle réside principalement dans les MLP et les embeddings de l'enseignant ; par conséquent, les taux d'apprentissage pour ces composants sont maintenus bas ou fixes pour éviter l'oubli catastrophique, tandis que le mélangeur de séquence est entraîné de manière plus agressive.
- Étape 3 : Extension de la longueur de contexte (Optionnel) : Le modèle est affiné sur des séquences plus longues (jusqu'à 16k jetons) en utilisant une perte d'entropie croisée standard sans modèle enseignant pour améliorer les capacités de contexte long. Alternativement, les auteurs proposent l'Étape 2a, qui étend la longueur de séquence à 4096 pendant la phase de distillation elle-même, rendant potentiellement inutile une étape 3 distincte.
2. Nouvelles Architectures
Les auteurs ont identifié que les architectures RWKV standard n'étaient pas toujours optimales pour la conversion. Ils ont introduit deux nouveaux variants :
- RAD-RWKV6 ("RADFinch") : Une modification de RWKV-6 utilisant un noyau d'attention linéaire à porte et des techniques d'équilibrage d'état pour améliorer la stabilité et l'ajustement.
- RAD-RWKV7 ("RADGoose") : Une modification de RWKV-7 qui supprime le mécanisme de "tokenshift" (qui n'apportait aucun bénéfice dans ce contexte) et applique directement des plongements positionnels rotatifs (RoPE). Cette architecture a démontré une convergence plus rapide et une perte de distillation plus faible par rapport aux versions non modifiées de RWKV-7.
3. Hyperparamètres et Données
- Jeu de données : Les auteurs se sont décantés pour DCLM (DataComp-LM) pour toutes les étapes de conversion, le jugeant supérieur à FineWeb ou FineWeb-Edu pour la conversion des modèles Qwen.
- Taux d'apprentissage : Un programme d'extinction cosinus (cosine annealing) est utilisé dans l'étape 1, commençant haut (10−3) pour aligner les états cachés et finissant proche du taux d'apprentissage final de pré-entraînement de l'enseignant (10−5). Les étapes 2 et 3 utilisent un taux d'apprentissage constant.
Principales Contributions
- Recette de distillation RADLADS : Un protocole détaillé, étape par étape, incluant des hyperparamètres spécifiques, des comptes de jetons et des choix de jeux de données qui permettent une conversion de haute qualité avec un minimum de données.
- Nouvelles Architectures : L'introduction de RAD-RWKV6 et RAD-RWKV7, qui sont optimisées pour le processus de conversion, offrant une inférence plus rapide et un meilleur alignement avec les modèles enseignants que leurs prédécesseurs non modifiés.
- Modèles à grande échelle : La conversion réussie de modèles open-source populaires Qwen2.5 en variantes d'attention linéaire à des échelles de 7B, 32B et 72B de paramètres.
- Libération en Open Source : La publication du code et des modèles convertis (QRWKV6/7-7B/32B/72B) sous licence Apache 2.0 (avec les restrictions de licence Qwen pour le modèle 72B), permettant à d'autres de répliquer le processus.
Résultats
Les modèles convertis atteignent des performances de pointe parmi les modèles récurrents purs de leur taille :
- Efficacité : Convertir un modèle de 72B coûte moins de 2 000 USD et nécessite environ 700 millions de jetons.
- Performance : Sur les benchmarks standards (Lambada, MMLU, ARC, etc.), les modèles RADLADS surpassent systématiquement les autres méthodes de conversion (ex: SUPRA, LoLCats, MOHAWK, ARWKV).
- Le QRWKV7-7B-Instruct atteint un score MMLU relatif de 92,4 % par rapport à son enseignant, surpassant nettement les autres conversions de 7B.
- Le QRWKV6-72B-Instruct atteint un score MMLU relatif de 89,9 %, établissant un nouveau SoTA pour les modèles de langage purement RNN à cette échelle.
- Vitesse d'inférence : Grâce au mécanisme d'attention linéaire, les modèles convertis montrent des accélérations significatives par rapport aux modèles enseignants à mesure que la longueur du contexte augmente. Par exemple, à 8k jetons d'entrée et 256 jetons de sortie, le modèle 32B QRWKV7 est 1,61x plus rapide que son enseignant Qwen3-32B ; à 6k d'entrée et 2k de sortie, l'accélération atteint 3,41x.
Limites et Revendications
Le papier reconnaît modestement plusieurs limites :
- Raisonnement et Contexte Long : Bien que les performances soient solides sur les benchmarks standards, les modèles présentent des limites dans les tâches de raisonnement complexe (ex: Minerva Math) et les tâches de contexte très long (ex: RULER), où ils ne s'améliorent pas aussi efficacement avec l'ajout de jetons de sortie que les modèles enseignants.
- Sensibilité de l'Architecture : Chaque nouvelle conception d'architecture nécessite des tests méticuleux pour assurer la compatibilité avec le protocole RADLADS. Par exemple, l'utilisation de GroupNorm/LayerNorm, communes dans certaines variantes d'attention linéaire, a provoqué une instabilité de l'entraînement à des échelles de 14B+ et a dû être remplacée par des techniques de pré-mise à l'échelle (pre-scaling).
- Alignement des Données : Les auteurs notent que les jeux de données de raisonnement peuvent nécessvoir un alignement plus étroit avec la distribution du modèle enseignant pour éviter les comportements de bouclage répétitifs, un défi qu'ils ont partiellement adressé en mélangeant DCLM avec OpenThoughts.
Signification
Le papier affirme que RADLADS offre une voie rentable pour démocratiser l'accès aux modèles d'attention linéaire à grande échelle. En réduissant le besoin d'entraînement de trillions de jetons à quelques centaines de millions, il permet aux chercheurs et aux organisations plus modestes de tester, entraîner et déployer de nouvelles variétés d'architectures RNN à grande échelle sans les coûts extrêmes associés au pré-entraînement. Les auteurs positionnent cela comme un outil pour accélérer le développement de la prochaine génération de variantes d'attention à état compressif.
Noyé(e) sous les articles dans votre domaine ?
Recevez des digests quotidiens des articles les plus récents correspondant à vos mots-clés de recherche — avec des résumés techniques, dans votre langue.
Recevez les meilleurs articles AI chaque semaine.
Adopté par des chercheurs de Stanford, Cambridge et de l'Académie des sciences.
Vérifiez votre boîte mail pour confirmer votre inscription.
Quelque chose s'est mal passé. Réessayer ?
Pas de spam, désinscription à tout moment.