如何在PyTorch中用RNNCell实现RNN?解决list无view属性报错
Let's break down and fix your code step by step. The core issue here is exactly what the error states: you're trying to call view() on a Python list (rnn_out), which doesn't support tensor operations like view(). There are also a few other small bugs we need to address to get the code running properly.
Key Fixes & Explanations
1. Convert the Hidden State List to a Tensor
Instead of storing each hidden state in a Python list and trying to call view() directly, use torch.stack() to convert the list of tensors into a single 3D tensor. You can then reshape this tensor to fit the linear layer's input requirements.
2. Fix Unreferenced Variables in forward()
In your forward method, you used hiddens instead of self.hiddens (since hiddens is an instance attribute of your SimpleRNN class). Also, initialize the hidden state on the same device as your input tensor (CPU/GPU) to avoid device mismatch errors later.
3. Correct Training Loop Variable Names
Your training loop references inputs and labels, but those variables don't exist in your code—you need to use x_one_hot and y_data instead. Additionally, CrossEntropyLoss expects the target tensor to be of type long, so we'll convert y_data to long() when calculating loss.
4. Align Input Shape with RNNCell Expectations
nn.RNNCell expects each time step input to be a batch of samples, so we need to adjust the input shape from [batch_size, seq_len, input_size] to [seq_len, batch_size, input_size] using permute(). This lets us iterate over each time step's batch of inputs correctly.
Corrected Full Code
import torch import torch.nn as nn torch.manual_seed(777) class SimpleRNN(nn.Module): def __init__(self, inputs, hiddens, n_class): super().__init__() self.rnn = nn.RNNCell(inputs, hiddens) self.linear = nn.Linear(hiddens, n_class) self.hiddens = hiddens def forward(self, x): # Initialize hidden state on the same device as input hx = torch.zeros((x.shape[1], self.hiddens), device=x.device) rnn_out = [] # Iterate over each time step (x has shape: [seq_len, batch_size, input_size]) for i in x: hx = self.rnn(i, hx) rnn_out.append(hx) # Convert list of tensors to a single tensor, then flatten rnn_out_tensor = torch.stack(rnn_out) # Shape: [seq_len, batch_size, hidden_size] # Reshape to (seq_len * batch_size, hidden_size) for linear layer linear_out = self.linear(rnn_out_tensor.view(-1, self.hiddens)) return linear_out # Hyperparameters seq_len = 6 # |hihell| == 6, equivalent to time step input_size = 5 # one-hot size batch_size = 1 # one sentence per batch num_classes = 5 # predicting 5 distinct character hidden_size = 4 # output from the RNN # Create an index to character mapping idx2char = ['h', 'i', 'e', 'l', 'o'] # Teach hihell -> ihello x_data = [[0, 1, 0, 2, 3, 3]] # hihell x_one_hot = [[[1, 0, 0, 0, 0], # h 0 [0, 1, 0, 0, 0], # i 1 [1, 0, 0, 0, 0], # h 0 [0, 0, 1, 0, 0], # e 2 [0, 0, 0, 1, 0], # l 3 [0, 0, 0, 1, 0]]] # l 3 x_one_hot = torch.Tensor(x_one_hot) # Adjust shape to [seq_len, batch_size, input_size] which RNNCell expects x_one_hot = x_one_hot.permute(1, 0, 2) y_data = torch.Tensor([1, 0, 2, 3, 3, 4]) # ihello rnn = SimpleRNN(input_size, hidden_size, num_classes) # Train the model num_epochs = 15 criterion = nn.CrossEntropyLoss() optimizer = torch.optim.Adam(rnn.parameters(), lr=0.1) for epoch in range(1, num_epochs + 1): optimizer.zero_grad() outputs = rnn(x_one_hot) # CrossEntropyLoss expects target to be long type loss = criterion(outputs, y_data.long()) loss.backward() optimizer.step() # Check the current predicted string _, idx = outputs.max(dim=1) idx = idx.data.numpy() result_str = [idx2char[c] for c in idx] print(f'epoch: {epoch}, loss: {loss.item():1.3f}') print('Predicted string: ', ''.join(result_str))
Additional Notes
- The
permute(1, 0, 2)call rearranges the input tensor dimensions so we can loop through each time step's batch of data, which is howRNNCellis designed to be used. - Initializing
hxwithdevice=x.deviceensures compatibility if you later move your model or tensors to a GPU, avoiding runtime device mismatch errors.
内容的提问来源于stack exchange,提问作者Eka

