import torch import torch.nn as nn import torch.nn.functional as F # Dummy data: batch size 1, 5 input steps, hidden size 4 encoder_outputs = torch.randn(1, 5, 4) # (batch, seq_len, hidden) decoder_hidden = torch.randn(1, 4) # (batch, hidden) ### Bahdanau Attention (additive) ### class BahdanauAttention(nn.Module): def __init__(self, hidden_size): super().__init__() self.W1 = nn.Linear(hidden_size, hidden_size) self.W2 = nn.Linear(hidden_size, hidden_size) self.v = nn.Linear(hidden_size, 1) def forward(self, encoder_outputs, decoder_hidden): seq_len = encoder_outputs.size(1) # Expand decoder hidden to shape (batch, seq_len, hidden) dec = decoder_hidden.unsqueeze(1).repeat(1, seq_len, 1) score = self.v( torch.tanh(self.W1(encoder_outputs) + self.W2(dec)) ).squeeze(-1) # (batch, seq_len) attn_weights = F.softmax(score, dim=1) context = torch.bmm(attn_weights.unsqueeze(1), encoder_outputs).squeeze(1) return attn_weights, context bahdanau = BahdanauAttention(hidden_size=4) bahdanau_weights, bahdanau_context = bahdanau(encoder_outputs, decoder_hidden) print("Bahdanau attention weights:", bahdanau_weights) print("Bahdanau context vector:", bahdanau_context) ### Luong Attention (dot) ### class LuongAttention(nn.Module): def __init__(self): super().__init__() def forward(self, encoder_outputs, decoder_hidden): # encoder_outputs: (batch, seq_len, hidden) # decoder_hidden: (batch, hidden) attn_weights = torch.bmm( encoder_outputs, decoder_hidden.unsqueeze(-1) ).squeeze(-1) # (batch, seq_len) attn_weights = F.softmax(attn_weights, dim=1) context = torch.bmm(attn_weights.unsqueeze(1), encoder_outputs).squeeze(1) return attn_weights, context luong = LuongAttention() luong_weights, luong_context = luong(encoder_outputs, decoder_hidden) print("Luong attention weights:", luong_weights) print("Luong context vector:", luong_context)