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

Google Colab中FUSE挂载海量嵌套数据集的文件存在性校验优化

优化Serengeti数据集文件存在性检查的方案

问题背景

已将Snapshot Serengeti数据集的云存储桶挂载到Google Colab环境,文件路径格式为/snapshotserengeti-unzipped/S1/B04/B04_R1/S1_B04_R1_PICT0003.JPG。现有两份CSV文件:一份记录所有图片路径与唯一capture_id,另一份关联capture_id与物种标注信息。

使用PyTorch构建SerengetiDataset类时,发现CSV中部分记录的文件在挂载目录中不存在,导致加载时报错。原解决方案是为CSV新增file_exists列标记文件存在性,但因数据集超700万张图片,且FUSE挂载的目录检查开销极大,单进程执行耗时超1小时。

现有Dataset实现代码:

class SerengetiDataset(data.Dataset):
    def __init__(self, root, csv_file1, csv_file2, transform=None):
        self.root = root
        self.transform = transform
        
        self.image_info = pd.read_csv(csv_file1)
        
        self.image_info['file_exists'] = self.image_info['image_path_rel'].apply(lambda x: os.path.exists(os.path.join(root, x)))
        
        self.labels_info = pd.read_csv(csv_file2)
        
        self.annotations = pd.merge(self.image_info, self.labels_info, on='capture_id')
        
        self.filenames = self.annotations['image_path_rel'].tolist()
        
    def __getitem__(self, index):
        filename = self.filenames[index]
        path = os.path.join(self.root, filename)
        image = Image.open(path).convert('RGB')
        
        label = self.annotations.loc[index, 'question__species']
        
        if self.transform is not None:
            image = self.transform(image)
            
        return image, label

原低效检查代码:

import pandas as pd

images_df = pd.read_csv('images.csv')

def file_exists(row):
    filename = os.path.join('/content/datasets/snapshotserengeti-unzipped', row['image_path_rel'])
    return os.path.exists(filename)

images_df['file_exists'] = images_df.apply(file_exists, axis=1)

images_df.to_csv('images_updated.csv', index=False)

优化方案

1. 批量遍历现有文件,集合匹配(推荐)

FUSE挂载下单个os.path.exists调用开销极高,改为一次性遍历挂载目录中所有存在的图片,将路径存入集合后与CSV匹配,集合查询为O(1)操作,整体效率大幅提升:

import os
import pandas as pd
from pathlib import Path

# 遍历挂载目录下所有JPG文件,生成相对路径集合
root_dir = "/content/datasets/snapshotserengeti-unzipped"
existing_paths = set()

# 使用Path.rglob高效递归遍历,仅匹配JPG文件
for img_path in Path(root_dir).rglob("*.JPG"):
    # 生成与CSV中image_path_rel格式一致的相对路径
    rel_path = os.path.relpath(img_path, root_dir)
    existing_paths.add(rel_path)

# 加载CSV并批量匹配存在性
images_df = pd.read_csv('images.csv')
images_df['file_exists'] = images_df['image_path_rel'].isin(existing_paths)

images_df.to_csv('images_updated.csv', index=False)

2. 多进程并行检查文件存在性

利用Colab的多CPU核心,通过多进程并行执行文件检查,大幅缩短单进程的等待时间:

import pandas as pd
import os
from torch.utils.data import Dataset
from multiprocessing import Pool

def check_exists(args):
    root, rel_path = args
    return os.path.exists(os.path.join(root, rel_path))

class SerengetiDataset(Dataset):
    def __init__(self, root, csv_file1, csv_file2, transform=None):
        self.root = root
        self.transform = transform
        
        self.image_info = pd.read_csv(csv_file1)
        self.labels_info = pd.read_csv(csv_file2)
        self.annotations = pd.merge(self.image_info, self.labels_info, on='capture_id')
        
        # 多进程并行检查,进程数设为CPU核心数
        with Pool(processes=os.cpu_count()) as pool:
            args_list = [(root, path) for path in self.annotations['image_path_rel']]
            exists_list = pool.map(check_exists, args_list)
        
        # 过滤无效条目并重置索引
        self.annotations = self.annotations[exists_list].reset_index(drop=True)
        self.filenames = self.annotations['image_path_rel'].tolist()
        
    def __getitem__(self, index):
        filename = self.filenames[index]
        path = os.path.join(self.root, filename)
        image = Image.open(path).convert('RGB')
        
        label = self.annotations.loc[index, 'question__species']
        
        if self.transform is not None:
            image = self.transform(image)
            
        return image, label

3. 懒加载+异常捕获(无需提前预处理)

若不想提前消耗时间预处理CSV,可在Dataset的__getitem__方法中捕获文件不存在的异常,自动跳过无效样本:

from torch.utils.data import Dataset
import pandas as pd
from PIL import Image

class SerengetiDataset(Dataset):
    def __init__(self, root, csv_file1, csv_file2, transform=None):
        self.root = root
        self.transform = transform
        
        self.image_info = pd.read_csv(csv_file1)
        self.labels_info = pd.read_csv(csv_file2)
        self.annotations = pd.merge(self.image_info, self.labels_info, on='capture_id').reset_index(drop=True)
        
    def __getitem__(self, index):
        while index < len(self.annotations):
            row = self.annotations.iloc[index]
            path = os.path.join(self.root, row['image_path_rel'])
            try:
                image = Image.open(path).convert('RGB')
                label = row['question__species']
                if self.transform is not None:
                    image = self.transform(image)
                return image, label
            except FileNotFoundError:
                # 跳过不存在的文件,索引自增
                index += 1
        raise IndexError("No valid samples available")
        
    def __len__(self):
        # 返回原始条目数,实际有效样本数需提前检查才能确定
        return len(self.annotations)

4. 直接调用云存储API查询(针对GCS/S3等云存储)

若挂载的是Google Cloud Storage(GCS)或AWS S3,直接调用云存储API查询文件存在性,绕过FUSE中间层,效率远高于文件系统调用:

from google.cloud import storage
import pandas as pd

# 初始化GCS客户端
client = storage.Client()
bucket_name = "your-bucket-name"
bucket = client.get_bucket(bucket_name)

images_df = pd.read_csv('images.csv')

def check_gcs_exists(rel_path):
    # 构建GCS对象路径,需与挂载的前缀一致
    blob_path = f"snapshotserengeti-unzipped/{rel_path}"
    return bucket.blob(blob_path).exists()

# 可结合多进程进一步加速
images_df['file_exists'] = images_df['image_path_rel'].apply(check_gcs_exists)
images_df.to_csv('images_updated.csv', index=False)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 21:22:02