基于LAD-GPT自定义数据集训练遇PyTorch IndexError问题求助
解决torch.nn.Embedding索引越界问题的排查与修改步骤
验证Token索引范围与Vocab大小的匹配性
尽管数据集有764个唯一token,仍需确认编码后的token索引最大值是否小于len(vocab):- 在数据加载或编码环节添加代码,打印所有样本的token索引最大值:
print(max(token_id for sample in dataset for token_id in sample['input_ids'])) - 检查Vocab构建是否包含所有必要的特殊token(如
<pad>、<unk>),且这些token已被计入len(vocab)。若编码时使用了未被纳入Vocab的特殊token,会直接导致索引越界。
- 在数据加载或编码环节添加代码,打印所有样本的token索引最大值:
确认Embedding层的初始化逻辑
排查nn.Embedding的初始化参数是否严格与Vocab长度一致:- 找到代码中
nn.Embedding的定义行,确保参数为nn.Embedding(len(vocab), embedding_dim),而非固定数值或被其他变量覆盖。 - 检查训练流程中是否存在Vocab实例不匹配的情况,比如构建Embedding时用的是旧Vocab,训练时加载了新的Vocab但未同步更新Embedding层参数。
- 找到代码中
排查数据预处理的Token映射逻辑
确保所有token都被正确映射到合法的Vocab索引:- 若Vocab包含
<unk>token,确认所有未登录词都被映射到<unk>的索引,而非直接生成超出范围的数字。 - 检查编码脚本的索引生成逻辑,避免出现索引从1开始但Vocab长度为764的情况(Embedding索引从0开始,最大合法索引为
len(vocab)-1)。
- 若Vocab包含
实时监控训练Batch的Token索引
在训练循环中添加临时打印代码,定位是否是特定Batch导致的问题:for batch in train_dataloader: curr_max = batch['input_ids'].max().item() if curr_max >= len(vocab): print(f"发现越界索引: {curr_max},Vocab长度: {len(vocab)}") # 后续训练代码
内容的提问来源于stack exchange,提问作者user9909570
相关产品推荐
相关产品推荐

