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

微调SAM模型时调用Dataset.from_dict加载数据集内核崩溃求助

解决Dataset.from_dict加载数据集时内核崩溃的问题

可能的原因及对应解决方案

  • 内存过载:

    • 检查data_dict里的图像/掩码是否是全量加载的numpy数组或PIL对象,这类数据会直接占满内存。改成存储文件路径,在Dataset的__getitem__里按需加载:
      # 原data_dict(可能导致内存爆炸)
      # data_dict = {"image": [np.array(img1), np.array(img2), ...], "mask": [np.array(mask1), ...]}
      
      # 修改后的data_dict(存路径)
      data_dict = {"image_path": ["path/to/img1.jpg", "path/to/img2.jpg", ...], "mask_path": ["path/to/mask1.png", ...]}
      
      # 自定义加载逻辑
      def load_data(examples):
          examples["image"] = np.array(Image.open(examples["image_path"]))
          examples["mask"] = np.array(Image.open(examples["mask_path"]))
          return examples
      
      dataset = Dataset.from_dict(data_dict).map(load_data)
      
    • 若必须提前加载,先取小样本测试,确认是否是内存问题:
      # 先测试10个样本
      small_data_dict = {k: v[:10] for k, v in data_dict.items()}
      dataset = Dataset.from_dict(small_data_dict)
      
  • 数据类型/格式不兼容:

    • SAM要求掩码是单通道二值数组(0/1),如果是多通道或RGB格式,会导致内存异常。转换掩码格式:
      def convert_mask(mask):
          if len(mask.shape) == 3:
              mask = np.mean(mask, axis=-1).astype(np.uint8)
          mask = (mask > 127).astype(np.uint8)  # 二值化处理
          return mask
      
      data_dict["train_masks"] = [convert_mask(m) for m in data_dict["train_masks"]]
      
    • 确保图像是uint8类型,避免float32等占用更多内存的格式:
      data_dict["train_images"] = [img.astype(np.uint8) for img in data_dict["train_images"]]
      
  • datasets库版本问题:

    • 旧版本可能存在内存泄漏或兼容性问题,升级到稳定版:
      pip install --upgrade datasets
      
    • 若升级后仍有问题,尝试回退到已知稳定版本,比如datasets==2.14.6:
      pip install datasets==2.14.6
      
  • 环境资源不足:

    • 关闭其他占用内存/显存的程序,确保当前环境有足够资源。GPU环境可通过nvidia-smi(Linux)或任务管理器(Windows)查看显存占用。
    • 云端环境(如Colab)可切换到更高配置的实例,提升显存容量。

调试建议

  • 执行前后打印内存使用情况,确认是否是内存溢出:
    import psutil
    
    def print_memory():
      process = psutil.Process()
      print(f"当前内存占用: {process.memory_info().rss / 1024 / 1024:.2f} MB")
    
    print_memory()
    dataset = Dataset.from_dict(data_dict)
    print_memory()
    
  • 逐样本检查数据格式,定位异常样本:
    for idx, (img, mask) in enumerate(zip(data_dict["train_images"], data_dict["train_masks"])):
        print(f"样本{idx} - 图像形状:{img.shape}, 类型:{img.dtype} | 掩码形状:{mask.shape}, 类型:{mask.dtype}")
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 17:35:03