You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中实现RNN(GRU)的序列生成推理?

How to Generate Sequences from a Trained GRU Language Model in PyTorch

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:

  1. Pass the current input and hidden state through the GRU
  2. Convert the model's output logits to a token (greedily or with sampling)
  3. Append the token to your generated sequence
  4. 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, and hidden_size).

内容的提问来源于stack exchange,提问作者Evan Pu

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.28 06:23:34