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

不使用nn.RNN搭建指定字符RNN时CrossEntropyLoss报错求助

手动实现字符RNN的RuntimeError排查与解决

我懂手动搭RNN踩坑的痛苦,不用nn.RNN确实容易在张量形状、维度匹配上栽跟头。结合你说的结构(输入到隐藏层Wxh、隐藏层循环Whh、隐藏到输出Who,Tanh隐藏层、Softmax输出层,用CrossEntropyLoss做损失),我整理了最容易触发RuntimeError的几个点和解决办法:

  • 张量维度不匹配(最常见!)
    手动计算时很容易搞错维度顺序,尤其是PyTorch默认的矩阵乘法是左乘,得仔细对齐:

    • 输入序列:假设你的输入是字符one-hot编码,形状通常是(seq_len, batch_size, vocab_size),那Wxh的形状应该是(hidden_size, vocab_size),这样torch.matmul(Wxh, x_t.T)(把x_t转成(vocab_size, batch_size))才能得到(hidden_size, batch_size)的输入特征,和隐藏层循环部分的输出维度一致。
    • 隐藏层循环:Whh的形状必须是(hidden_size, hidden_size),保证torch.matmul(Whh, h_prev)的结果和输入特征维度匹配,相加后过Tanh得到当前隐藏状态h_t。
    • 输出层:Who的形状是(vocab_size, hidden_size),计算出的logits是(vocab_size, batch_size),但CrossEntropyLoss要求输入是(batch_size, vocab_size),所以必须转置一下,比如logits = logits.transpose(0, 1)。
  • CrossEntropyLoss的输入误区
    这个坑很多人都会踩:

    • CrossEntropyLoss自带Softmax计算,你不需要手动在输出层加Softmax!如果手动过了Softmax再喂给损失函数,会导致数值不稳定或者维度/类型不匹配的错误,直接把Who @ h_t的logits结果丢进去就行。
    • 目标标签的形状:如果你的目标是(seq_len, batch_size)的整数标签,那输入logits要对应调整成(seq_len*batch_size, vocab_size),目标标签要展平成(seq_len*batch_size,);或者按时间步逐个计算损失,每个时间步的logits是(batch_size, vocab_size),目标是(batch_size,),这样就不会报错。
  • 权重初始化的形状错误
    检查你的权重矩阵形状:

    • Wxh:输入特征数是vocab_size,输出是hidden_size → (hidden_size, vocab_size)
    • Whh:隐藏层到自身的循环权重 → (hidden_size, hidden_size)
    • Who:隐藏层到输出层,输出对应每个字符的logits → (vocab_size, hidden_size)
      另外别忘了加偏置项bh((hidden_size, 1))和bo((vocab_size, 1)),很多人会漏掉偏置导致维度或者数值错误。
  • 初始隐藏状态的形状问题
    初始隐藏状态h0的形状必须和后续时间步的h_t一致,比如批量训练时,h0应该是(hidden_size, batch_size),不能是(hidden_size,)(单样本才可以)。可以用torch.zeros(hidden_size, batch_size, device=device)来初始化,确保和输入在同一设备上(CPU/GPU)。

给你贴一段关键部分的正确示例代码,参考一下:

import torch

# 参数设置
vocab_size = 100
hidden_size = 128
seq_len = 20
batch_size = 32
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 初始化权重和偏置
Wxh = torch.randn(hidden_size, vocab_size, requires_grad=True, device=device)
Whh = torch.randn(hidden_size, hidden_size, requires_grad=True, device=device)
Who = torch.randn(vocab_size, hidden_size, requires_grad=True, device=device)
bh = torch.randn(hidden_size, 1, requires_grad=True, device=device)
bo = torch.randn(vocab_size, 1, requires_grad=True, device=device)

# 模拟输入和目标
x = torch.randn(seq_len, batch_size, vocab_size, device=device)  # (seq_len, batch_size, vocab_size)
targets = torch.randint(0, vocab_size, (seq_len, batch_size), device=device)  # (seq_len, batch_size)

# 初始化隐藏状态
h_prev = torch.zeros(hidden_size, batch_size, device=device)
loss_fn = torch.nn.CrossEntropyLoss()
total_loss = 0.0

for t in range(seq_len):
    x_t = x[t].T  # 转置为(vocab_size, batch_size)
    # 计算当前隐藏状态
    h_t = torch.tanh(torch.matmul(Wxh, x_t) + torch.matmul(Whh, h_prev) + bh)
    # 计算输出logits
    logits = torch.matmul(Who, h_t) + bo  # (vocab_size, batch_size)
    logits = logits.T  # 转为(batch_size, vocab_size),符合Loss要求
    # 累加损失
    total_loss += loss_fn(logits, targets[t])

# 反向传播更新参数
total_loss.backward()

如果还是报错,你可以把具体的错误提示贴出来(比如维度不匹配的具体数值),这样能更精准定位问题~

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 08:14:29