如何在PyTorch中实现RNN(GRU)的序列生成推理?
Hey there! I totally get where you're coming from—training a GRU with teacher forcing is pretty straightforward, but switching over to generating sequences from scratch can feel like a bit of a leap at first. Let's walk through exactly how to make this work in PyTorch, no fancy new APIs required.
The core idea here is that during inference, you can't rely on the ground truth tokens (like you did with teacher forcing). Instead, you'll build the sequence one token at a time: taking the model's output from the previous step, turning it into a token, and feeding that right back in as the next input. Here's how to implement this:
Step 1: Set Up for Inference
First, make sure your model is in evaluation mode (this turns off dropout/batch norm behaviors that only make sense during training) and moved to the right device. You'll also start with your <start> token and the initial all-zero hidden state, just like you did during training.
import torch # Assume your trained model is loaded and ready device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) model.eval() # Critical for proper inference behavior # Define your <start> token index (match your vocabulary) start_token_idx = 0 # Replace with your actual <start> token ID current_input = torch.tensor([[start_token_idx]], device=device) # Shape: (batch_size=1, seq_len=1) # Initialize hidden state (match your GRU's hidden size and layer count) hidden_size = model.gru.hidden_size h_n = torch.zeros(1, 1, hidden_size, device=device) # Shape: (num_layers, batch_size, hidden_size)
Step 2: Run the Autoregressive Generation Loop
Now you'll loop for as many steps as you want to generate. In each iteration:
- Pass the current input and hidden state through the GRU
- Convert the model's output logits to a token (greedily or with sampling)
- Append the token to your generated sequence
- Update the input and hidden state for the next step
Here's the code for greedy generation (picking the most likely token each time):
generated_sequence = [start_token_idx] max_generated_length = 10 # Adjust to how many tokens you want for _ in range(max_generated_length): # Disable gradient computation to save memory and speed things up with torch.no_grad(): # Forward pass through GRU gru_output, h_n = model.gru(current_input, h_n) # Pass GRU output through your final linear layer to get vocab logits logits = model.fc(gru_output) # Shape: (1, 1, vocab_size) # Pick the token with the highest probability (greedy search) predicted_token_idx = torch.argmax(logits, dim=-1).item() # Add to our sequence generated_sequence.append(predicted_token_idx) # Set this token as the input for the next iteration current_input = torch.tensor([[predicted_token_idx]], device=device) # Convert indices back to actual tokens using your vocabulary # generated_tokens = [your_vocab[idx] for idx in generated_sequence] # print(generated_tokens)
Step 3: Optional: Add Sampling for More Natural Outputs
Greedy search can lead to repetitive, boring sequences. If you want more diverse results, try top-k sampling with temperature scaling to adjust how "random" the outputs are:
def sample_token(logits, temperature=1.0, top_k=50): # Temperature adjusts the sharpness of the probability distribution scaled_logits = logits / temperature # Keep only the top-k most likely tokens to avoid rare, nonsensical ones top_k_logits, top_k_indices = torch.topk(scaled_logits, top_k, dim=-1) # Convert logits to probabilities probs = torch.softmax(top_k_logits, dim=-1) # Sample a token from the probability distribution sampled_pos = torch.multinomial(probs, num_samples=1).item() # Map back to the original vocabulary index return top_k_indices[0, 0, sampled_pos].item() # Replace the argmax line in the loop with this: predicted_token_idx = sample_token(logits, temperature=0.8, top_k=30)
Do You Need a Different API?
Short answer: No. PyTorch's GRU module doesn't have a built-in "roll out" function, but you don't need one. All you have to do is manually handle the iterative sequence building, reusing the hidden state from each step and feeding the model's output back in as the next input. This is standard practice for autoregressive generation with RNNs/GRUs/LSTMs.
Quick Pro Tips
- Always use
model.eval()during generation—forgetting this can lead to weird, inconsistent outputs from dropout layers. torch.no_grad()is essential here; it stops PyTorch from tracking gradients, which saves a ton of memory and speeds up inference.- Double-check that your input and hidden state dimensions match what your GRU expects (pay attention to
num_layers,batch_size, andhidden_size).
内容的提问来源于stack exchange,提问作者Evan Pu

