Transformer decoder

Modelli Transformer con PyTorch

James Chapman

Curriculum Manager, DataCamp

Dall'originale al transformer solo decoder

Architettura originale del transformer

Modelli Transformer con PyTorch

Dall'originale al transformer solo decoder

Architettura transformer solo decoder

Generazione autoregressiva di sequenze: generazione e completamento di testo

Modelli Transformer con PyTorch

Dall'originale al transformer solo decoder

Architettura transformer solo decoder

Generazione autoregressiva di sequenze: generazione e completamento di testo

Self-attention mascherata multi-head

  • Nasconde i token successivi nella sequenza
Modelli Transformer con PyTorch

Dall'originale al transformer solo decoder

Architettura transformer solo decoder

Generazione autoregressiva di sequenze: generazione e completamento di testo

Self-attention mascherata multi-head

  • Nasconde i token successivi nella sequenza

Testa del transformer solo decoder

  • Linear + Softmax sul vocabolario
  • Predice i token successivi più probabili
Modelli Transformer con PyTorch

Self-attention mascherata/attenzione causale

Self-attention mascherata

  • Chiave del comportamento autoregressivo o causale
  • Maschera di attenzione triangolare (causale)
Modelli Transformer con PyTorch

Self-attention mascherata/attenzione causale

Self-attention mascherata

  • Chiave del comportamento autoregressivo o causale
  • Maschera di attenzione triangolare (causale)
  • Ogni token guarda solo ai token precedenti nella sequenza
Modelli Transformer con PyTorch

Self-attention mascherata/attenzione causale

Self-attention mascherata

tgt_mask = (1 - torch.triu(
  torch.ones(1, seq_len, seq_len), diagonal=1)
).bool()
  • Chiave del comportamento autoregressivo o causale
  • Maschera di attenzione triangolare (causale)
  • Ogni token guarda solo ai token precedenti nella sequenza

    • "favorite": "orange", "is", "my", "favorite"
  • Attenzione causale forzata: predice la prossima parola da generare, ad es. "fruit"

Modelli Transformer con PyTorch

Layer del decoder

class DecoderLayer(nn.Module):
    def __init__(self, d_model, num_heads, d_ff, dropout):
        super().__init__()
        self.self_attn = MultiHeadAttention(d_model, num_heads)
        self.ff_sublayer = FeedForwardSubLayer(d_model, d_ff)
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, tgt_mask):
        attn_output = self.self_attn(x, x, x, tgt_mask)
        x = self.norm1(x + self.dropout(attn_output))
        ff_output = self.ff_sublayer(x)
        x = self.norm2(x + self.dropout(ff_output))
        return x
Modelli Transformer con PyTorch

Corpo e testa del transformer decoder

class TransformerDecoder(nn.Module):
    def __init__(self, vocab_size, d_model, num_layers, num_heads, d_ff, dropout, max_seq_length):
        super(TransformerDecoder, self).__init__()
        self.embedding = InputEmbeddings(vocab_size, d_model)
        self.positional_encoding = PositionalEncoding(d_model, max_seq_length)
        self.layers = nn.ModuleList([DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)])

self.fc = nn.Linear(d_model, vocab_size)
def forward(self, x, tgt_mask): x = self.embedding(x) x = self.positional_encoding(x) for layer in self.layers: x = layer(x, tgt_mask)
x = self.fc(x) return F.log_softmax(x, dim=-1)
  • self.fc: layer lineare di output con vocab_size neuroni
  • Aggiungi self.fc e attivazione softmax nel forward pass
Modelli Transformer con PyTorch

Istanziamento del transformer solo decoder

decoder = TransformerDecoder(vocab_size, d_model, num_layers, num_heads, d_ff, dropout, max_seq_length=seq_length)

output = decoder(input_sequence, tgt_mask)
tensor([[[ -9.4692,  -9.8429,  -9.3077,  ...,  -9.9523, -10.2669,  -9.7084],
         [ -9.1556,  -9.6133, -10.0923,  ...,  -9.3810,  -9.0420,  -9.1780],
         ...,
         [ -9.5327, -10.3534,  -9.8443,  ...,  -9.8170,  -8.8491,  -8.8322],
         [ -9.6086,  -9.6336, -10.1595,  ...,  -9.8550,  -9.9955,  -8.7121]],

        [[ -9.5865,  -8.0360,  -8.5056,  ...,  -9.9855,  -9.5677,  -9.0352],
         [ -9.7213,  -8.6451,  -8.3779,  ...,  -9.2994,  -9.2601,  -9.8509],
         ...,
         [ -9.0471,  -9.7410, -10.0160,  ..., -10.0195,  -9.4651,  -8.9605],
         [ -9.5767, -10.2692,  -8.8394,  ...,  -8.3458,  -9.1479, -10.0650]]],
       grad_fn=<LogSoftmaxBackward0>)
Modelli Transformer con PyTorch

Esercitiamoci!

Modelli Transformer con PyTorch

Preparing Video For Download...