Files
temperature-based-fertility…/code/models/lstm.py
T
Alex Blank c6defa2065 fixes
2025-05-19 13:59:16 +02:00

47 lines
1.5 KiB
Python

from torch import nn
class LSTMModel(nn.Module):
def __init__(self,
input_dim: int,
output_dim: int,
cnn_channels: int,
cnn_kernel_size: int,
embed_dim: int,
lstm_hidden_size=64,
num_layers=1):
super().__init__()
self.cnn = nn.Sequential(
nn.Conv1d(input_dim, cnn_channels, kernel_size=cnn_kernel_size, padding=1),
nn.BatchNorm1d(cnn_channels),
nn.ReLU(),
nn.Conv1d(cnn_channels, embed_dim, kernel_size=cnn_kernel_size, padding=1),
nn.BatchNorm1d(embed_dim),
nn.ReLU()
)
self.lstm = nn.LSTM(
input_size=embed_dim,
hidden_size=lstm_hidden_size,
num_layers=num_layers,
batch_first=True,
)
self.head = nn.Sequential(
nn.Linear(lstm_hidden_size, output_dim) # Output is a scalar Δt
)
def forward(self, x):
# x: (batch_size, seq_len, input_size)
x = x.permute(0, 2, 1)
# x: (batch_size, input_size, seq_len)
x = self.cnn(x)
# x: (batch_size, embed_dim, seq_len)
x = x.permute(0, 2, 1)
# x: (batch_size, seq_len, embed_dim)
x, _ = self.lstm(x)
# x: (batch_size, seq_len, lstm_hidden_size)
x = x[:, -1, :] # Get the last time step
# x: (batch_size, lstm_hidden_size)
x = self.head(x)
# x: (batch_size, output_size)
return x