class BahdanauAttention([nn.Module](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module")): def __init__(self, hidden_size): super([BahdanauAttention](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module"), self).__init__() self.Wa = [nn.Linear](https://docs.pytorch.org/docs/stable/generated/torch.nn.Linear.html#torch.nn.Linear "torch.nn.Linear")(hidden_size, hidden_size) self.Ua = [nn.Linear](https://docs.pytorch.org/docs/stable/generated/torch.nn.Linear.html#torch.nn.Linear "torch.nn.Linear")(hidden_size, hidden_size) self.Va = [nn.Linear](https://docs.pytorch.org/docs/stable/generated/torch.nn.Linear.html#torch.nn.Linear "torch.nn.Linear")(hidden_size, 1) def forward(self, query, keys): scores = self.Va([torch.tanh](https://docs.pytorch.org/docs/stable/generated/torch.tanh.html#torch.tanh "torch.tanh")(self.Wa(query) + self.Ua(keys))) scores = scores.squeeze(2).unsqueeze(1) weights = [F.softmax](https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.softmax.html#torch.nn.functional.softmax "torch.nn.functional.softmax")(scores, dim=-1) context = [torch.bmm](https://docs.pytorch.org/docs/stable/generated/torch.bmm.html#torch.bmm "torch.bmm")(weights, keys) return context, weights class AttnDecoderRNN([nn.Module](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module")): def __init__(self, hidden_size, output_size, dropout_p=0.1): super([AttnDecoderRNN](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module"), self).__init__() self.embedding = [nn.Embedding](https://docs.pytorch.org/docs/stable/generated/torch.nn.Embedding.html#torch.nn.Embedding "torch.nn.Embedding")(output_size, hidden_size) self.attention = [BahdanauAttention](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module")(hidden_size) self.gru = [nn.GRU](https://docs.pytorch.org/docs/stable/generated/torch.nn.GRU.html#torch.nn.GRU "torch.nn.GRU")(2 * hidden_size, hidden_size, batch_first=True) self.out = [nn.Linear](https://docs.pytorch.org/docs/stable/generated/torch.nn.Linear.html#torch.nn.Linear "torch.nn.Linear")(hidden_size, output_size) self.dropout = [nn.Dropout](https://docs.pytorch.org/docs/stable/generated/torch.nn.Dropout.html#torch.nn.Dropout "torch.nn.Dropout")(dropout_p) def forward(self, encoder_outputs, encoder_hidden, target_tensor=None): batch_size = encoder_outputs.size(0) decoder_input = [torch.empty](https://docs.pytorch.org/docs/stable/generated/torch.empty.html#torch.empty "torch.empty")(batch_size, 1, dtype=[torch.long](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype"), [device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")=[device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")).fill_(SOS_token) decoder_hidden = encoder_hidden decoder_outputs = [] attentions = [] for i in range(MAX_LENGTH): decoder_output, decoder_hidden, attn_weights = self.forward_step( decoder_input, decoder_hidden, encoder_outputs ) decoder_outputs.append(decoder_output) attentions.append(attn_weights) if target_tensor is not None: # Teacher forcing: Feed the target as the next input decoder_input = target_tensor[:, i].unsqueeze(1) # Teacher forcing else: # Without teacher forcing: use its own predictions as the next input _, topi = decoder_output.topk(1) decoder_input = topi.squeeze(-1).detach() # detach from history as input decoder_outputs = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")(decoder_outputs, dim=1) decoder_outputs = [F.log_softmax](https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.log_softmax.html#torch.nn.functional.log_softmax "torch.nn.functional.log_softmax")(decoder_outputs, dim=-1) attentions = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")(attentions, dim=1) return decoder_outputs, decoder_hidden, attentions def forward_step(self, input, hidden, encoder_outputs): embedded = self.dropout(self.embedding(input)) query = hidden.permute(1, 0, 2) context, attn_weights = self.attention(query, encoder_outputs) input_gru = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")((embedded, context), dim=2) output, hidden = self.gru(input_gru, hidden) output = self.out(output) return output, hidden, attn_weights