Aller au contenu principal
Expert Cybersécurité & IAv9.0
Centres de ressources conformité
Besoin d'un accompagnement expert ?
Devis personnalisé sous 24h — audit, conformité, incident
Checklists Sécurité — Audit & Durcissement
Formats disponibles
📄 PDF 📊 Excel 🌐 Web

11 checklists professionnelles couvrant 2 200+ points de contrôle. Téléchargement gratuit, aucune inscription.

Weight Tying

ia

Dé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

Un projet cybersécurité ?

Expert dispo · Réponse 24h

Devis