darshan204/HumanValuesUncover
02
1import torch2import torch.nn as nn3 4class BiLSTMClassifier(nn.Module):5 def __init__(self, embedding_dim, hidden_dim, output_dim, n_layers, bidirectional, dropout):6 super().__init__()7 self.embedding_dim = embedding_dim8 self.hidden_dim = hidden_dim9 self.output_dim = output_dim10 self.n_layers = n_layers11 self.bidirectional = bidirectional12 self.dropout = dropout13 14 self.lstm = nn.LSTM(embedding_dim,15 hidden_dim,16 num_layers=n_layers,17 bidirectional=bidirectional,18 dropout=dropout,19 batch_first=True)20 21 self.fc = nn.Linear(hidden_dim * 2 if bidirectional else hidden_dim, output_dim)22 self.dropout = nn.Dropout(dropout)23 24 def forward(self, text, text_lengths):25 # text = [batch size, sent len]26 27 # pack sequence28 packed_embedded = nn.utils.rnn.pack_padded_sequence(text, text_lengths.cpu(), batch_first=True, enforce_sorted=False)29 30 packed_output, (hidden, cell) = self.lstm(packed_embedded)31 32 # unpack sequence33 output, output_lengths = nn.utils.rnn.pad_packed_sequence(packed_output, batch_first=True)34 35 # output = [batch size, sent len, hidden dim * n directions]36 # hidden = [n layers * n directions, batch size, hidden dim]37 38 # concat the final forward (hidden[-2,:,:]) and backward (hidden[-1,:,:]) hidden layers39 # and apply dropout40 hidden = self.dropout(torch.cat((hidden[-2,:,:], hidden[-1,:,:]), dim=1) if self.bidirectional else hidden[-1,:,:])41 42 # hidden = [batch size, hidden dim * n directions]43 44 return self.fc(hidden)