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

训练PyTorch RNN时损失不下降,求排查(附自定义情感分析RNN代码)

Troubleshooting Your PyTorch RNN's Stagnant Loss for Sentiment Analysis

Let's walk through the key issues in your code that are likely causing the loss to not decrease, plus some extra tips to get your model training properly.

Critical Bugs in Your Forward Pass

First, the biggest red flag is in your forward method:

hidden_new = self.relu(self.i2h(input)+self.h2h(hidden))
output = self.h2o(hidden)

You're calculating an updated hidden state (hidden_new) but then using the old, unupdated hidden state to compute the output. That means your model isn't actually learning from the current input at all—no wonder the loss isn't moving! Fix this by using hidden_new for the output layer:

output = self.h2o(hidden_new)

Confusing Activation Function Naming

You assigned nn.Tanh() to self.relu—this is not just bad readability, it could lead to bugs later if you forget this mismatch. Rename it to something accurate:

self.tanh = nn.Tanh()
# Then in forward:
hidden_new = self.tanh(self.i2h(input) + self.h2h(hidden))

Suboptimal Output Activation for Sentiment Analysis

For sentiment analysis (usually binary or multi-class classification):

  • If it's binary classification, skip manually adding nn.LogSigmoid() and use BCEWithLogitsLoss instead. This loss function integrates the sigmoid activation internally, which is more numerically stable than applying it separately.
  • If it's multi-class, use CrossEntropyLoss (which combines Softmax and log loss) and remove the LogSigmoid entirely.

Missing Hidden State Initialization

Did you remember to initialize the hidden state at the start of each sequence or batch? If you're feeding in sequences without resetting/initializing hidden, your model will carry over stale state from previous samples, which breaks training. Add a helper method to initialize it:

def init_hidden(self, batch_size=1):
    return torch.zeros(batch_size, self.hidden_size).to(input.device)  # Match your device

Call this at the start of processing each batch or sequence.

Additional Troubleshooting Tips

  • Weight Initialization: Default linear layer weights might lead to gradient vanishing/explosion in RNNs. Try initializing weights with Xavier uniform initialization:
    nn.init.xavier_uniform_(self.i2h.weight)
    nn.init.xavier_uniform_(self.h2h.weight)
    nn.init.xavier_uniform_(self.h2o.weight)
    
  • Optimizer Choice: If you're using SGD, switch to Adam—it's generally better for RNNs and adapts learning rates automatically. Start with a learning rate of 1e-3.
  • Data Preprocessing: Double-check that your text data is correctly converted to numerical embeddings (e.g., using PyTorch's Embedding layer), and that input dimensions match your input_size. If you're using padded sequences, use pack_padded_sequence to ignore padding during training.

Fixed Example Code

Here's a cleaned-up version of your RNN with the above fixes:

import torch
import torch.nn as nn

class RNN(nn.Module):
    def __init__(self, input_size, hidden_size, output_size):
        super().__init__()
        self.hidden_size = hidden_size
        
        # Linear layers
        self.i2h = nn.Linear(input_size, hidden_size)
        self.h2h = nn.Linear(hidden_size, hidden_size)
        self.h2o = nn.Linear(hidden_size, output_size)
        
        # Activation
        self.tanh = nn.Tanh()
        
        # Initialize weights
        nn.init.xavier_uniform_(self.i2h.weight)
        nn.init.xavier_uniform_(self.h2h.weight)
        nn.init.xavier_uniform_(self.h2o.weight)

    def forward(self, input, hidden):
        hidden_new = self.tanh(self.i2h(input) + self.h2h(hidden))
        output = self.h2o(hidden_new)
        return output, hidden_new

    def init_hidden(self, batch_size=1, device="cpu"):
        return torch.zeros(batch_size, self.hidden_size).to(device)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 08:04:48