class AttentionLayer(nn.Module): def __init__(self, input_dim, attention_dim=64): super(SimpleAttention, self).__init__() # Linear layers for attention computation self.query = nn.Linear(input_dim, attention_dim) self.key = nn.Linear(input_dim, attention_dim) self.value = nn.Linear(input_dim, attention_dim) # Scaling factor self.scale = torch.sqrt(torch.FloatTensor([attention_dim])) def forward(self, x): # x shape: (batch_size, sequence_length, input_dim) batch_size, seq_len, input_dim = x.size() # Compute Q, K, V Q = self.query(x) # (batch_size, seq_len, attention_dim) K = self.key(x) # (batch_size, seq_len, attention_dim) V = self.value(x) # (batch_size, seq_len, attention_dim) attention_scores = torch.matmul(Q, K.transpose(-2, -1)) / self.scale # Scaled dot-product attention attention_weights = F.softmax(attention_scores, dim=-1) # Convert attention weights to probabilities attended_output = torch.matmul(attention_weights, V) # Apply attention to values return attended_output, attention_weights class TransformerBlock(nn.Module): """ A single transformer block composed of self-attention and a feed-forward network. """ def __init__(self, embed_dim, ffn_hidden_dim): """ Args: embed_dim (int): The dimensionality of the model's embeddings. ffn_hidden_dim (int): The dimensionality of the hidden layer in the FFN. """ super(TransformerBlock, self).__init__() self.attention = SimpleAttention(embed_dim, embed_dim) self.norm1 = nn.LayerNorm(embed_dim) self.norm2 = nn.LayerNorm(embed_dim) self.ffn = nn.Sequential( nn.Linear(embed_dim, ffn_hidden_dim), nn.ReLU(), nn.Linear(ffn_hidden_dim, embed_dim) ) def forward(self, x): """ Forward pass for the transformer block. Args: x (torch.Tensor): Input tensor of shape (batch_size, sequence_length, embed_dim). Returns: torch.Tensor: The output tensor of the transformer block. """ # Self-attention part attended, _ = self.attention(x) # Add & Norm (residual connection) x = self.norm1(attended + x) # Feed-forward part ffn_out = self.ffn(x) # Add & Norm (residual connection) x = self.norm2(ffn_out + x) return x class TransformerEncoder(nn.Module): """ A transformer encoder that stacks multiple TransformerBlocks. """ def __init__(self, num_layers, embed_dim, ffn_hidden_dim, seq_len, output_dim): """ Args: num_layers (int): The number of transformer blocks to stack. embed_dim (int): The dimensionality of the model's embeddings. ffn_hidden_dim (int): The dimensionality of the hidden layer in the FFN. seq_len (int): The length of the input sequences. output_dim (int): The dimensionality of the final output (e.g., number of classes). """ super(TransformerEncoder, self).__init__() # Create a list of transformer blocks self.layers = nn.ModuleList( [TransformerBlock(embed_dim, ffn_hidden_dim) for _ in range(num_layers)] ) # Final classification head self.classifier = nn.Linear(embed_dim * seq_len, output_dim) def forward(self, x): """ Forward pass for the full transformer encoder. Args: x (torch.Tensor): Input tensor of shape (batch_size, sequence_length, embed_dim). Returns: torch.Tensor: The final output logits from the classifier. """ # Pass input through all transformer blocks for layer in self.layers: x = layer(x) # Flatten the output for the classifier x = x.view(x.size(0), -1) # Final classification output = self.classifier(x) return output