smartmcq-demo / model.py
balamaniansp's picture
custom bidirectional gru mode
87546cb verified
Raw
History Blame Contribute Delete
3.01 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from huggingface_hub import PyTorchModelHubMixin
class CustomBiGRUMCQModel(nn.Module, PyTorchModelHubMixin):
"""
Bidirectional GRU model with custom global attention pooling
"""
#vocab_size, embedding_dim=256, hidden_dim=256, num_layers=2, dropout=0.1
def __init__(self, config:dict):
super().__init__()
self.config = config
self.vocab_size = config.get('vocab_size', 30000)
self.embedding_dim = config.get('embedding_dim',256)
self.hidden_dim = config.get('hidden_dim',256)
self.num_layers = config.get('num_layers',2)
self.embedding = nn.Embedding(self.vocab_size, self.embedding_dim, padding_idx=0)
self.embedding_dropout = nn.Dropout(config.get('dropout', 0.1))
self.gru = nn.GRU(
input_size = self.embedding_dim,
hidden_size = self.hidden_dim,
num_layers = self.num_layers,
bidirectional=True,
batch_first=True,
dropout = config.get('dropout',0.1) if self.num_layers > 1 else 0
)
# Dot product attention pooling head
self.attention_query = nn.Linear(self.hidden_dim * 2, 1, bias=False)
#Final classification layer
self.classifier = nn.Linear(self.hidden_dim*2, 1)
self.dropout = nn.Dropout(config.get('dropout',0.1))
def forward(self, input_ids, attention_mask=None, label=None, **kwargs):
input_idx = input_ids
#expected shape [batch_size, num_of_choice=5, seq_length]
batch_size, num_of_choices, seq_length = input_idx.shape
flat_input_idx = input_ids.view(batch_size * num_of_choices, seq_length)
embedded = self.embedding(flat_input_idx)
embedded = self.embedding_dropout(embedded)
gru_out, _ = self.gru(embedded) # Shape: [batch_size * 5, seq_length, hidden_dim * 2]
attentions_scores = self.attention_query(gru_out).squeeze(-1)
# Adapt Attentions mask to the choice dimension
flaten_mask = None
if attention_mask is not None:
flat_mask = attention_mask.view(batch_size * num_of_choices, seq_length)
attentions_scores = attentions_scores.masked_fill(flat_mask == 0, -1e9)
attention_weight = F.softmax(attentions_scores, dim=-1).unsqueeze(-1)
pooled_output = torch.sum(gru_out * attention_weight, dim=1)
pooled_output = self.dropout(pooled_output)
raw_scores = self.classifier(pooled_output).squeeze(-1) # Shape: [batch_size * 5]
reshape_logits = raw_scores.view(batch_size, num_of_choices) # Shape: [batch_size, 5]
loss = None
if label is not None:
loss_fct = nn.CrossEntropyLoss()
loss = loss_fct(reshape_logits, label)
return {
'loss': loss,
'logits': reshape_logits
}