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

基于LAD-GPT自定义数据集训练遇PyTorch IndexError问题求助

解决torch.nn.Embedding索引越界问题的排查与修改步骤
  • 验证Token索引范围与Vocab大小的匹配性
    尽管数据集有764个唯一token,仍需确认编码后的token索引最大值是否小于len(vocab):

    1. 在数据加载或编码环节添加代码,打印所有样本的token索引最大值:print(max(token_id for sample in dataset for token_id in sample['input_ids']))
    2. 检查Vocab构建是否包含所有必要的特殊token(如<pad>、<unk>),且这些token已被计入len(vocab)。若编码时使用了未被纳入Vocab的特殊token,会直接导致索引越界。
  • 确认Embedding层的初始化逻辑
    排查nn.Embedding的初始化参数是否严格与Vocab长度一致:

    1. 找到代码中nn.Embedding的定义行,确保参数为nn.Embedding(len(vocab), embedding_dim),而非固定数值或被其他变量覆盖。
    2. 检查训练流程中是否存在Vocab实例不匹配的情况,比如构建Embedding时用的是旧Vocab,训练时加载了新的Vocab但未同步更新Embedding层参数。
  • 排查数据预处理的Token映射逻辑
    确保所有token都被正确映射到合法的Vocab索引:

    1. 若Vocab包含<unk>token,确认所有未登录词都被映射到<unk>的索引,而非直接生成超出范围的数字。
    2. 检查编码脚本的索引生成逻辑,避免出现索引从1开始但Vocab长度为764的情况(Embedding索引从0开始,最大合法索引为len(vocab)-1)。
  • 实时监控训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 13:39:58