import torch import torch.nn as nn class SimpleRNN(nn.Module): """ Simple RNN model for text classification. Architecture: Embedding -> RNN -> Hidden Layers -> Output """ def __init__( self, vocab_size: int, embedding_dim: int = 128, hidden_dim: int = 256, num_layers: int = 2, num_classes: int = 2, dropout: float = 0.3, num_hidden_nodes: int = 128 ): """ Args: vocab_size: Size of vocabulary embedding_dim: Dimension of word embeddings hidden_dim: Hidden dimension of RNN num_layers: Number of RNN layers num_classes: Number of output classes dropout: Dropout rate num_hidden_nodes: Number of nodes in hidden layer """ super(SimpleRNN, self).__init__() self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=0) self.rnn = nn.RNN( input_size=embedding_dim, hidden_size=hidden_dim, num_layers=num_layers, batch_first=True, dropout=dropout if num_layers > 1 else 0 ) # Hidden layer (flexible nodes) self.hidden_layer = nn.Sequential( nn.Linear(hidden_dim, num_hidden_nodes), nn.ReLU(), nn.Dropout(dropout) ) # Output layer self.output_layer = nn.Linear(num_hidden_nodes, num_classes) def forward(self, x): """ Forward pass. Args: x: Input tensor of shape (batch_size, seq_length) Returns: Output tensor of shape (batch_size, num_classes) """ # Embedding embedded = self.embedding(x) # (batch_size, seq_length, embedding_dim) # RNN rnn_out, hidden = self.rnn(embedded) # rnn_out: (batch_size, seq_length, hidden_dim) # Use the last output of the sequence # For simple RNN, this is the final hidden state after processing all tokens last_output = rnn_out[:, -1, :] # (batch_size, hidden_dim) # Hidden layer hidden_out = self.hidden_layer(last_output) # (batch_size, num_hidden_nodes) # Output layer output = self.output_layer(hidden_out) # (batch_size, num_classes) return output