Embedding索引越界异常排查:输入最大值合规仍报错
排查
Embedding Index out of range错误的思路 核心矛盾分析
从你提供的信息来看:
- 输入张量的最大值为
14,而self.src_word_emb是Embedding(15, 512, padding_idx=1)——正常情况下,Embedding的合法索引范围是0到num_embeddings-1(即0-14),输入完全在合规范围内,理论上不应触发索引越界错误。
具体排查步骤
- 检查是否误用了解码器的Embedding层
你提到“编码器嵌入仅接受小于embedding_dim_decoder的输入”,很大概率是代码中误将编码器输入传给了解码器的Embedding层。比如变量名写错,用self.tgt_word_emb(解码器嵌入)处理编码器输入,而解码器的num_embeddings可能小于15,直接导致索引越界。 - 确认输入张量的真实数据类型与值
有时打印的张量可能存在截断显示,或隐式类型转换(比如从int64转为uint8),可以执行src_seq.unique()输出所有唯一值,确认是否存在超出0-14的数值;同时检查src_seq.dtype,确保是整数类型(如torch.long)。 - 验证Embedding层的初始化与赋值
确认self.src_word_emb是否为你打印的那个Embedding实例,排查代码其他位置是否有重新赋值或替换操作,比如初始化后又执行self.src_word_emb = nn.Embedding(xxx, 512)覆盖了原有层。 - 排查输入的动态修改逻辑
查看emb = self.src_word_emb(src_seq)之前的代码,确认是否对src_seq做过偏移、索引映射等修改,导致输入值超出Embedding的合法范围。
验证代码示例
可以在报错前添加以下代码做快速验证:
# 验证输入范围 print("输入最大值:", src_seq.max().item()) print("输入最小值:", src_seq.min().item()) print("所有唯一值:", src_seq.unique()) # 验证Embedding层参数 print("Embedding num_embeddings:", self.src_word_emb.num_embeddings) print("Embedding padding_idx:", self.src_word_emb.padding_idx)
内容的提问来源于stack exchange,提问作者Malte
相关产品推荐
相关产品推荐

