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
相关产品推荐
相关产品推荐

