datadabllp

concept of self-attention

Aug 26th, 2024
683
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 1.57 KB | None | 0 0
  1. import torch
  2. import torch.nn as nn
  3.  
  4. class SelfAttention(nn.Module):
  5.     def __init__(self, embed_size, heads):
  6.         super(SelfAttention, self).__init__()
  7.         self.embed_size = embed_size
  8.         self.heads = heads
  9.         self.head_dim = embed_size // heads
  10.  
  11.         self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)
  12.         self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)
  13.         self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)
  14.         self.fc_out = nn.Linear(heads * self.head_dim, embed_size)
  15.  
  16.     def forward(self, values, keys, query, mask):
  17.         N = query.shape[0]
  18.         value_len, key_len, query_len = values.shape[1], keys.shape[1], query.shape[1]
  19.  
  20.         # Split embedding into self.heads pieces
  21.         values = values.reshape(N, value_len, self.heads, self.head_dim)
  22.         keys = keys.reshape(N, key_len, self.heads, self.head_dim)
  23.         queries = query.reshape(N, query_len, self.heads, self.head_dim)
  24.  
  25.         values = self.values(values)
  26.         keys = self.keys(keys)
  27.         queries = self.queries(queries)
  28.  
  29.         # Scaled dot-product attention
  30.         energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])
  31.         if mask is not None:
  32.             energy = energy.masked_fill(mask == 0, float("-1e20"))
  33.  
  34.         attention = torch.softmax(energy / (self.embed_size ** (1/2)), dim=3)
  35.  
  36.         out = torch.einsum("nhql,nlhd->nqhd", [attention, values]).reshape(
  37.             N, query_len, self.heads * self.head_dim
  38.         )
  39.  
  40.         out = self.fc_out(out)
  41.         return out
  42.  
Advertisement
Add Comment
Please, Sign In to add comment