To be specific, if I swap out my RNN above with the following model, it gives me 90% accuracy on the same data set with all the same hyperparameters.
class SWEM(nn.Module):
def __init__(self, vocab_size, embedding_size, hidden_dim, num_outputs):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embedding_size)
self.fc1 = nn.Linear(embedding_size, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, num_outputs)
def forward(self, x):
embed = self.embedding(x)
embed_mean = torch.mean(embed, dim=0)
h = self.fc1(embed_mean)
h = torch.nn.functional.relu(h)
h = self.fc2(h)
return h
So this, to me, implies one of two things: either the RNN code I have has a bug, or it can’t cope with the input data somehow. The input x consists of padded sequences, collated in this way by the DataLoader:
def collator(batch):
labels = torch.tensor([example[0] for example in batch])
sentences = [example[1] for example in batch]
data = pad_sequence(sentences)
return [data, labels]
The data itself is the ag_news dataset in CSV form, looking like this: AG News Classification Dataset | Kaggle