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

PyTorch DataLoader报错:'KeyError'对象不可迭代,求解决方案

PyTorch DataLoader TypeError 问题排查与修复

问题背景

运行PyTorch训练代码时触发错误:

TypeError: Caught TypeError in DataLoader worker process 0. TypeError: 'KeyError' object is not iterable

已尝试为DataLoader添加drop_last=True参数,且确认torch/torchvision版本兼容、CUDA可用,但问题未解决。

错误根源

问题出在自定义UdacitySelfDrivingDataset类的__getitem__方法中:当传入的idx不在DataFrame索引中时,代码返回了一个KeyError对象,而不是抛出异常。DataLoader的worker进程会尝试迭代这个返回值(期望得到(img, target)二元组),但KeyError对象不可迭代,因此触发TypeError。

原错误代码片段:

def __getitem__(self, idx):
    
    if idx in self.df.index:
        row = self.df.loc[[idx]]
    else:
        return KeyError(f"Element {idx} not in dataframe")

修复方案

将返回KeyError对象改为抛出KeyError异常,这是PyTorch Dataset类的标准行为。修改后的__getitem__方法如下:

def __getitem__(self, idx):
    
    if idx in self.df.index:
        row = self.df.loc[[idx]]
    else:
        raise KeyError(f"Element {idx} not in dataframe")
    
    # 后续代码保持不变
    img_path = os.path.join(self.root, "images", row['frame'].iloc[0])
    img = Image.open(img_path).convert("RGB")
    
    # 过滤宽高为0的无效框
    h = row['ymax'] - row['ymin']
    w = row['xmax'] - row['xmin']
    filter_idx = (h > 0) & (w > 0)
    row = row[filter_idx]
    
    # 提取边界框坐标
    boxes = row[['xmin', 'ymin', 'xmax', 'ymax']].values
    boxes = torch.as_tensor(boxes, dtype=torch.float32)
    
    # 提取标签
    labels = torch.as_tensor(row['class_id'].values, dtype=int)

    image_id = torch.tensor([idx])
    area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])
    
    # 假设没有拥挤区域
    iscrowd = torch.zeros((row.shape[0],), dtype=torch.int64)

    target = {}
    target["boxes"] = boxes
    target["labels"] = labels
    target["image_id"] = image_id
    target["area"] = area
    target["iscrowd"] = iscrowd

    if self.transform is not None:
        img, target = self.transform(img, target)

    return img, target

额外验证点

  • 确认get_data_loaders中创建的sampler使用的索引范围不超过Dataset的__len__()返回值(当前代码中train_idx和valid_idx从torch.randperm(n_tot)生成,n_tot = len(train_data),这部分逻辑正确)。
  • 若使用了thinning参数,确认Dataset初始化后self.n_images与实际保留的图片数量一致,避免sampler生成超出范围的索引。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 08:59:17