Spaces:
Running on Zero
Running on Zero
Download models/abstractive/decoder.py from fady21/bilingual-summarizer-api: direct link, hf CLI and curl.
- Browser
- Download file 3.24 kB
-
https://huggingface.co/spaces/fady21/bilingual-summarizer-api/resolve/main/models/abstractive/decoder.py
- Command line
-
hf download hf://spaces/fady21/bilingual-summarizer-api/models/abstractive/decoder.py
-
curl -L -o decoder.py https://huggingface.co/spaces/fady21/bilingual-summarizer-api/resolve/main/models/abstractive/decoder.py
3.24 kB
| import torch | |
| import torch.nn as nn | |
| from typing import Tuple | |
| from .attention import BahdanauAttention | |
| class Seq2SeqDecoder(nn.Module): | |
| """ | |
| Attention-guided GRU Decoder for Seq2Seq Abstractive Summarization. | |
| """ | |
| def __init__( | |
| self, | |
| output_dim: int, | |
| emb_dim: int, | |
| enc_hid_dim: int, | |
| dec_hid_dim: int, | |
| dropout: float = 0.2, | |
| pad_idx: int = 0 | |
| ): | |
| super().__init__() | |
| self.output_dim = output_dim | |
| self.emb_dim = emb_dim | |
| self.enc_hid_dim = enc_hid_dim | |
| self.dec_hid_dim = dec_hid_dim | |
| self.attention = BahdanauAttention(enc_hid_dim * 2, dec_hid_dim) | |
| self.embedding = nn.Embedding(output_dim, emb_dim, padding_idx=pad_idx) | |
| self.dropout = nn.Dropout(dropout) | |
| # GRU takes concatenated [embedded token, context vector] | |
| # context vector size is enc_hid_dim * 2 | |
| self.rnn = nn.GRU((enc_hid_dim * 2) + emb_dim, dec_hid_dim, batch_first=True) | |
| # Projection layer to vocabulary logits | |
| self.fc_out = nn.Linear((enc_hid_dim * 2) + dec_hid_dim + emb_dim, output_dim) | |
| def forward( | |
| self, | |
| input_token: torch.Tensor, | |
| hidden: torch.Tensor, | |
| encoder_outputs: torch.Tensor, | |
| mask: torch.Tensor = None | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """ | |
| Args: | |
| input_token: [batch_size] (current step token index) | |
| hidden: [batch_size, dec_hid_dim] (previous decoder hidden state) | |
| encoder_outputs: [batch_size, src_len, enc_hid_dim * 2] | |
| mask: [batch_size, src_len] | |
| Returns: | |
| prediction: [batch_size, output_dim] (vocabulary probability distribution) | |
| hidden: [batch_size, dec_hid_dim] (updated decoder hidden state) | |
| a: [batch_size, src_len] (attention weights) | |
| """ | |
| # input_token: [batch_size, 1] | |
| input_token = input_token.unsqueeze(1) | |
| embedded = self.dropout(self.embedding(input_token)) # [batch_size, 1, emb_dim] | |
| # Calculate attention weights: [batch_size, src_len] | |
| a = self.attention(hidden, encoder_outputs, mask=mask) | |
| # Compute context vector: weighted sum of encoder outputs | |
| # a.unsqueeze(1): [batch_size, 1, src_len] | |
| # encoder_outputs: [batch_size, src_len, enc_hid_dim * 2] | |
| # context: [batch_size, 1, enc_hid_dim * 2] | |
| context = torch.bmm(a.unsqueeze(1), encoder_outputs) | |
| # Combine embedded input and context vector | |
| rnn_input = torch.cat((embedded, context), dim=2) # [batch_size, 1, emb_dim + enc_hid_dim * 2] | |
| # Step RNN: output shape [batch_size, 1, dec_hid_dim] | |
| output, hidden = self.rnn(rnn_input, hidden.unsqueeze(0)) | |
| hidden = hidden.squeeze(0) # [batch_size, dec_hid_dim] | |
| # Combine output, context, and embedding for final prediction | |
| output = output.squeeze(1) # [batch_size, dec_hid_dim] | |
| context = context.squeeze(1) # [batch_size, enc_hid_dim * 2] | |
| embedded = embedded.squeeze(1) # [batch_size, emb_dim] | |
| prediction = self.fc_out(torch.cat((output, context, embedded), dim=1)) # [batch_size, output_dim] | |
| return prediction, hidden, a | |