DL_Empa / model.py
Daimka's picture
Upload 3 files
607293d verified
Raw
History Blame Contribute Delete
2.41 kB
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