[lstm](https://docs.pytorch.org/docs/stable/generated/torch.nn.LSTM.html#torch.nn.LSTM "torch.nn.LSTM") = [nn.LSTM](https://docs.pytorch.org/docs/stable/generated/torch.nn.LSTM.html#torch.nn.LSTM "torch.nn.LSTM")(3, 3) # Input dim is 3, output dim is 3 [inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [[torch.randn](https://docs.pytorch.org/docs/stable/generated/torch.randn.html#torch.randn "torch.randn")(1, 3) for _ in range(5)] # make a sequence of length 5 # initialize the hidden state. hidden = ([torch.randn](https://docs.pytorch.org/docs/stable/generated/torch.randn.html#torch.randn "torch.randn")(1, 1, 3), [torch.randn](https://docs.pytorch.org/docs/stable/generated/torch.randn.html#torch.randn "torch.randn")(1, 1, 3)) for [i](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") in [inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"): # Step through the sequence one element at a time. # after each step, hidden contains the hidden state. [out](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), hidden = [lstm](https://docs.pytorch.org/docs/stable/generated/torch.nn.LSTM.html#torch.nn.LSTM "torch.nn.LSTM")([i](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").view(1, 1, -1), hidden) # alternatively, we can do the entire sequence all at once. # the first value returned by LSTM is all of the hidden states throughout # the sequence. the second is just the most recent hidden state # (compare the last slice of "out" with "hidden" below, they are the same) # The reason for this is that: # "out" will give you access to all hidden states in the sequence # "hidden" will allow you to continue the sequence and backpropagate, # by passing it as an argument to the lstm at a later time # Add the extra 2nd dimension [inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")).view(len([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")), 1, -1) hidden = ([torch.randn](https://docs.pytorch.org/docs/stable/generated/torch.randn.html#torch.randn "torch.randn")(1, 1, 3), [torch.randn](https://docs.pytorch.org/docs/stable/generated/torch.randn.html#torch.randn "torch.randn")(1, 1, 3)) # clean out hidden state [out](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), hidden = [lstm](https://docs.pytorch.org/docs/stable/generated/torch.nn.LSTM.html#torch.nn.LSTM "torch.nn.LSTM")([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), hidden) print([out](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) print(hidden)