Kom igångKom igång gratis

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

Visa kurs

Övningsinstruktioner

  • Definiera ett LSTM-lager i metoden .__init__() och tilldela det till self.lstm.
  • Initiera det första långtidsminnet, det dolda tillståndet c0, med nollor i metoden forward().
  • Skicka alla tre indata till LSTM-lagret i metoden forward(): 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
Redigera och kör kod