import torch import torch.nn as nn # Define constants n_heads = 3 query_key_dim = 64 value_dim = 64 class MultiHeadSelfAttention(nn.Module): def __init__(self): super(MultiHeadSelfAttention, self).__init__() # Defining the linear layers that construct the query, key, and value self.W_Q = nn.Linear(d_model, query_key_dim * n_heads) # Projects input to [batch x sequence x (q/k_dim*num_heads)] self.W_K = nn.Linear(d_model, query_key_dim * n_heads) # Projects input to [batch x sequence x (q/k_dim*num_heads)] self.W_V = nn.Linear(d_model, value_dim * n_heads) # Projects input to [batch x sequence x (v_dim*num_heads)] self.dot_prod_attn = ScaledDotProductAttention() # Parameterless system that calculates attention self.proj_back = nn.Linear(value_dim * n_heads, d_model) # Projects final output of mhsa back into model dimension def forward(self, embedding): # passing embedding through dense networks qs = self.W_Q(embedding) # [batch_size x sequence_len x (query_key_dim * n_heads)] ks = self.W_K(embedding) # [batch_size x sequence_len x (query_key_dim * n_heads)] vs = self.W_V(embedding) # [batch_size x sequence_len x (value_dim * n_heads)] #dividing out heads #[batch_size, sequence_len, q/k/v_dim, n_heads] qs = qs.view(batch_size, max_input_length, query_key_dim, n_heads) ks = ks.view(batch_size, max_input_length, query_key_dim, n_heads) vs = vs.view(batch_size, max_input_length, value_dim, n_heads) #moving the head dimension next to the batch dimension #[batch_size x n_heads x sequence_len x q/k/v_dim] qs = qs.permute(0, 3, 1, 2) ks = ks.permute(0, 3, 1, 2) vs = vs.permute(0, 3, 1, 2) #combining batch and head dimension #[batch_size*n_heads x sequence_len x q/k/v_dim] qs = qs.reshape(-1, max_input_length, query_key_dim) ks = ks.reshape(-1, max_input_length, query_key_dim) vs = vs.reshape(-1, max_input_length, value_dim) #passing batches/heads of self attention through attn #[batch_size*n_heads x sequence_len x q/k/v_dim] head_results, _ = self.dot_prod_attn(qs,ks,vs) #seperating heads #[batch_size x n_heads x sequence_len x v_dim] head_results = head_results.reshape(batch_size,n_heads,max_input_length,value_dim) #moving the head dimension to the end #[batch_size x sequence_len x query_key_dim x n_heads] head_results = head_results.permute(0, 2, 3, 1) #combining the last dim to effectively concatonate the result of the heads #[batch_size x sequence_len x query_key_dim*n_heads] head_results = head_results.reshape(batch_size, max_input_length, -1) #projecting result of head back into model dimension return self.proj_back(head_results) # Example usage sample_embeddings = torch.tensor([[[1.1] * d_model] * max_input_length] * batch_size).to(device) print("Sample embeddings shape:", sample_embeddings.shape) sample = MultiHeadSelfAttention().to(device) output = sample(sample_embeddings) print('Output shape of mhsa:', output.shape)