Views
No views yet
1
2class RNN(nn.Module):
3 def __init__(self):
4 super().__init__()
5 num_classes = 7
6 hidden_size = 128
7 dropout = 0.4
8 embedding_dim = 768
9 num_layers = 1 # Increase the number of layers to 3
10 self.rnn = nn.GRU(embedding_dim, hidden_size, num_layers, batch_first=True, dropout =dropout)
11 self.dropout = nn.Dropout(dropout)
12 self.fc1 = nn.Linear(hidden_size, num_classes)
13
14 def forward(self, x):
15 mean_x = torch.mean(x, dim=1, keepdim=True)
16 batch_size, num_models, seq_len, hidden_size = mean_x.shape
17 x = mean_x.reshape(batch_size*num_models, seq_len, hidden_size)
18
19 x, _ = self.rnn(x)
20
21 x = F.relu(x)
22 x = F.max_pool1d(x.transpose(1, 2), x.size(1)).squeeze(2)
23
24 x = self.dropout(x)
25 logit = self.fc1(x)
26 return logit