Weight Tying
iaDéfinition
Le Weight Tying (partage de poids) est une technique d'optimisation architecturale des LLMs qui consiste a partager la matrice de la couche d'embedding d'entree (token_embedding) et la matrice de la tete de prediction de sortie (lm_head), en utilisant litteralement la meme matrice de poids pour les deux. Cette technique, introduite par Press et Wolf (2017), reduit significativement le nombre de parametres sans degrader les performances. La justification theorique : les deux matrices ont la meme dimension — (vocab_size, d_model) pour l'embedding et (d_model, vocab_size) pour le lm_head (transposee). L'intuition est que les deux matrices representent la meme information : une correspondance entre les tokens du vocabulaire et leur representation dans l'espace d'embedding. Il est coherent que le vecteur qui 'encode' le token A soit similaire au vecteur qui 'predit' le token A comme output. Le gain en nombre de parametres : pour LLaMA 3.1 8B (vocab_size=128K, d_model=4096), la matrice d'embedding contient 128000 * 4096 = 524M parametres, soit 6.5% des 8B parametres totaux. Sans weight tying, ces 524M parametres sont comptabilises deux fois (input + output) — le weight tying les partage et economise 524M parametres. LLaMA 3 utilise le weight tying, comme la plupart des LLMs modernes. Les implications pour l'inference et le fine-tuning : lors du fine-tuning avec LoRA, les adaptateurs sur lm_head et sur l'embedding sont necessaires meme avec weight tying (si le tokenizer est etendu avec de nouveaux tokens). Si le vocabulary est fixe, le weight tying reduit aussi la memoire VRAM lors du fine-tuning car une seule matrice est stockee. Les cas ou le weight tying est deconseille : (1) Extension du vocabulaire (nouveaux tokens specialises) — l'embedding et la tete de prediction ont des roles differents pour les nouveaux tokens, (2) Modeles encoder-decoder (T5, BART) — l'encoder et le decoder utilisent des espaces semantiques differents, rendant le weight tying sous-optimal, (3) Modeles tres profonds (>100B) ou la separation peut permettre une meilleure specialisation.
Implementation Weight Tying en PyTorch
import torch.nn as nn
class LLM(nn.Module):
def __init__(self, vocab_size=128000, d_model=4096):
super().__init__()
self.embed_tokens = nn.Embedding(vocab_size, d_model)
self.transformer_blocks = nn.ModuleList([...]) # Couches Transformer
self.norm = nn.RMSNorm(d_model)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
# Weight Tying : partager les poids embed_tokens et lm_head
self.lm_head.weight = self.embed_tokens.weight # Meme objet tenseur
def forward(self, input_ids):
x = self.embed_tokens(input_ids) # Utilise embed_tokens.weight
for block in self.transformer_blocks:
x = block(x)
x = self.norm(x)
logits = self.lm_head(x) # Utilise le MEME weight (transposee)
return logits
Articles liés
Expert en cybersécurité offensive et intelligence artificielle. Pentest, audit et développement IA sur-mesure.
Services
- Audit Infrastructure
- Audit Kubernetes
- Audit Microsoft 365
- Audit Sécurité Réseau
- Analyse de Risques
- Audit Active Directory
- Audit Application Web
- Audit Cloud (AWS/Azure/GCP)
- Audit Messagerie
- Audit API (OWASP Top 10)
- Audit DevSecOps & CI/CD
- Audit Code Source (SAST)
- Audit Postes de Travail
- Audit Sauvegarde & Résilience
- Audit OT/SCADA (IEC 62443)
- Développement IA
- Formations
Ressources
Projets & Outils
© 2026 Ayi NEDJIMI Consultants. Tous droits réservés.
Un projet cybersécurité ?
Expert dispo · Réponse 24h