Kom igångKom igång gratis

Skapa positionskodningar

Att bädda in tokens är en bra start, men inbäddningarna saknar fortfarande information om varje tokens position i sekvensen. För att åtgärda detta använder transformerarkitekturen positionskodningar. Dessa kodar in positionsinformation från varje token i inbäddningarna.

Du ska skapa en PositionalEncoding-klass med följande parametrar:

  • d_model: dimensionaliteten hos indatainbäddningarna
  • max_seq_length: den maximala sekvenslängden (eller sekvenslängden om alla sekvenser har samma längd)

Den här övningen är en del av kursen

Transformermodeller med PyTorch

Visa kurs

Övningsinstruktioner

  • Skapa en nollmatris med dimensionerna max_seq_length gånger d_model.
  • Utför sinus- och cosinusberäkningarna på position * div_term för att skapa de jämna och udda positionsinbäddningsvärdena.
  • Se till att pe inte är en inlärningsbar parameter under träningen.
  • Lägg till de transformerade positionsinbäddningarna till indatans tokeninbäddningar, x.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

class PositionalEncoding(nn.Module):
    def __init__(self, d_model, max_seq_length):
        super().__init__()
        # Create a matrix of zeros of dimensions max_seq_length by d_model
        pe = ____
        position = torch.arange(0, max_seq_length, dtype=torch.float).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model))
        
        # Perform the sine and cosine calculations
        pe[:, 0::2] = torch.____(position * div_term)
        pe[:, 1::2] = torch.____(position * div_term)
        # Ensure pe isn't a learnable parameter during training
        self.____('____', pe.unsqueeze(0))
        
    def forward(self, x):
        # Add the positional embeddings to the token embeddings
        return ____ + ____[:, :x.size(1)]

pos_encoding_layer = PositionalEncoding(d_model=512, max_seq_length=4)
output = pos_encoding_layer(token_embeddings)
print(output.shape)
print(output[0][0][:10])
Redigera och kör kod