不使用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)。
- 输入序列:假设你的输入是字符one-hot编码,形状通常是
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,),这样就不会报错。
- CrossEntropyLoss自带Softmax计算,你不需要手动在输出层加Softmax!如果手动过了Softmax再喂给损失函数,会导致数值不稳定或者维度/类型不匹配的错误,直接把
权重初始化的形状错误
检查你的权重矩阵形状:- 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)),很多人会漏掉偏置导致维度或者数值错误。
- Wxh:输入特征数是vocab_size,输出是hidden_size →
初始隐藏状态的形状问题
初始隐藏状态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
相关产品推荐
相关产品推荐

