LSTM-nätverk
Som du redan vet används vanliga RNN-celler sällan i praktiken. Ett mer vanligt förekommande alternativ, som hanterar långa sekvenser betydligt bättre, är Long Short-Term Memory-celler – eller LSTM. I den här övningen bygger du ett eget LSTM-nätverk!
Den viktigaste implementationsskillnaden jämfört med det RNN-nätverk du byggde tidigare är att LSTM-celler har två dolda tillstånd i stället för ett. Det innebär att du behöver initiera det extra dolda tillståndet och skicka det vidare till LSTM-cellen.
torch och torch.nn har redan importerats, så du kan börja koda direkt!
Den här övningen är en del av kursen
Fördjupad djupinlärning med PyTorch
Övningsinstruktioner
- Definiera ett LSTM-lager i metoden
.__init__()och tilldela det tillself.lstm. - Initiera det första långtidsminnet, det dolda tillståndet
c0, med nollor i metodenforward(). - Skicka alla tre indata till
LSTM-lagret i metodenforward(): indata för det aktuella tidssteget samt en tupel med de två dolda tillstånden.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
class Net(nn.Module):
def __init__(self, input_size):
super().__init__()
# Define lstm layer
____ = ____(
input_size=1,
hidden_size=32,
num_layers=2,
batch_first=True,
)
self.fc = nn.Linear(32, 1)
def forward(self, x):
h0 = torch.zeros(2, x.size(0), 32)
# Initialize long-term memory
c0 = ____
# Pass all inputs to lstm layer
out, _ = ____
out = self.fc(out[:, -1, :])
return out