class WordPredictionLSTMModel(nn.Module): def __init__(self, num_embed, embed_dim, pad_idx, lstm_hidden_dim, lstm_num_layers, output_dim, dropout): super().__init__() self.vocab_size = num_embed self.embed = nn.Embedding(num_embed, embed_dim, pad_idx) self.lstm = nn.LSTM(embed_dim, lstm_hidden_dim, lstm_num_layers, batch_first=True, dropout=dropout) self.fc = nn.Sequential( nn.Linear(lstm_hidden_dim, lstm_hidden_dim * 4), nn.LayerNorm(lstm_hidden_dim * 4), nn.LeakyReLU(), nn.Dropout(p=dropout), nn.Linear(lstm_hidden_dim * 4, output_dim), ) # def forward(self, x): x = self.embed(x) x, _ = self.lstm(x) x = self.fc(x) x = x.permute(0, 2, 1) return x # #