47 lines
1.5 KiB
Python
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
|