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

PyTorch使用Subset构建DataLoader时触发IndexError问题排查

解决PyTorch Subset + DataLoader训练时的IndexError问题

问题背景

在主动学习循环中,每次通过采样算法获取新索引,用这些索引构建torch.utils.data.Subset,再创建DataLoader用于训练,但运行时持续抛出IndexError: list index out of range错误。

报错原因分析

从报错栈可以定位到,错误发生在DataLoader迭代读取Subset数据时:Subset.__getitem__尝试用self.indices[idx]访问原数据集,但该索引超出了原数据集的有效范围(即self.indices[idx] >= len(dataset)或为负数)。

常见触发场景:

  • 采样算法返回无效索引:ALGO1/ALGO2输出的new_idx1/new_idx2包含超出原数据集范围的数值(如负数、大于等于len(dataset)的数)。
  • 局部索引未映射为全局索引:若采样算法是基于remaining_data(剩余样本的Subset)运行,返回的是该Subset的局部索引(0到len(remaining_data)-1),直接将其当作原数据集的全局索引使用会导致错误。
  • 集合操作意外引入无效值:虽然remaining_idx从range(len(dataset))生成,但如果采样算法返回了不在remaining_idx中的值,后续集合操作可能保留无效索引。

解决方案

1. 校验并过滤无效索引

在合并新索引到训练集前,添加校验逻辑,确保所有索引都是原数据集的有效索引:

# 对采样得到的新索引进行有效性过滤
def filter_valid_indices(indices, dataset_len, existing_idx):
    return [idx for idx in indices if 0 <= idx < dataset_len and idx not in existing_idx]

# 在final_loop的采样步骤后添加过滤:
dataset_len = len(dataset)
new_idx1 = filter_valid_indices(new_idx1, dataset_len, train_idx)
new_idx2 = filter_valid_indices(new_idx2, dataset_len, train_idx)

2. 正确映射局部索引到全局索引

如果采样算法是基于剩余样本的Subset运行,返回的是局部索引,需要将其映射为原数据集的全局索引:

# 构建剩余样本的Subset和对应的索引列表
remaining_data = Subset(dataset, remaining_idx)
# 假设ALGO1接收remaining_data的输入,返回局部索引(如0, 5, 10)
local_idx1 = ALGO1(...)  # 输出局部索引列表
# 映射为原数据集的全局索引
new_idx1 = [remaining_idx[idx] for idx in local_idx1]

# ALGO2同理
local_idx2 = ALGO2(...)
new_idx2 = [remaining_idx[idx] for idx in local_idx2]

3. 检查采样算法的输入输出

确认ALGO1和ALGO2的输入参数是否正确(原代码中inputs变量未定义,属于笔误),确保算法输出的索引符合预期(是原数据集的全局索引,或是需要映射的局部索引)。

修正后的关键代码片段

def final_loop(model, dataset, val_data, test_data, budget, gamma = 0.5, rounds = 45, model_name = "engrad.pt", init_samples = 1000,keep_old = True):
    device = torch.device("cuda")
    init_idx = random.sample(range(0,len(dataset)), init_samples)
    train_idx = init_idx
    train_step_data = Subset(dataset, train_idx)
    train_loader = DataLoader(train_step_data, batch_size = 20, shuffle = True)
    remaining_idx = list(set(range(0,len(dataset))) - set(train_idx))
    valid_loader = DataLoader(val_data, batch_size = 20)
    test_loader = DataLoader(test_data, batch_size = 20)
    test_acc = 0
    test_acc_list = [test_cifar(model, test_loader, device = "cuda")]
    validation = []
    samples = [len(train_idx)]
    print("test acc", test_acc_list[0])
    
    dataset_len = len(dataset)
    def filter_valid_indices(indices):
        return [idx for idx in indices if 0 <= idx < dataset_len and idx not in train_idx]

    for i in range(rounds):
        print("rounds = ",i+1,"------", "Datapoints = ", len(train_step_data))
        train_loss, val_loss = train_cifar(train_loader, valid_loader, model, epochs = 1, criterion= criterion, device = "cuda", model_name = model_name, save = True)
        validation.extend(val_loss)
        test_acc = test_cifar(model, test_loader, device = "cuda")
        test_acc_list.append(test_acc)

        # 修正:先构建剩余数据集,获取正确输入
        remaining_data = Subset(dataset, remaining_idx)
        # 示例:根据实际逻辑从remaining_data提取输入喂给ALGO
        # inputs = ...  

        # Sampling method1
        local_idx1 = ALGO1(...)  # 返回remaining_data的局部索引
        new_idx1 = [remaining_idx[idx] for idx in local_idx1]
        new_idx1 = filter_valid_indices(new_idx1)
        remaining_idx = list(set(remaining_idx) - set(new_idx1))

        # Sampling method2
        print("running loss_dep")
        local_idx2 = ALGO2(...)  # 返回remaining_data的局部索引
        new_idx2 = [remaining_idx[idx] for idx in local_idx2]
        new_idx2 = filter_valid_indices(new_idx2)
        remaining_idx = list(set(remaining_idx) - set(new_idx2))
        
        train_idx = list(set(train_idx + new_idx1 + new_idx2))
        train_step_data = Subset(dataset,train_idx)

        print("New data points selected")
        samples.append(len(train_idx))
        train_loader = DataLoader(train_step_data, batch_size = 20, shuffle = True)
        model.load_state_dict(torch.load("/content/" + model_name))
        print("Best model loaded")
        print("test acc",test_acc)
    return test_acc_list, validation, samples

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 14:45:37