修改PyTorch DataLoader,避免跨目录混合图像序列批次
解决方案:按目录独立生成图像批次
要实现同目录图像独立生成批次、跨目录不凑批次、不足批次的目录/剩余图像直接丢弃的需求,我们需要从两个核心部分修改代码:重构Dataset以按目录分组图像,以及自定义BatchSampler来生成符合要求的批次索引。
1. 重构Dataset:按目录分组存储图像
首先修改MainDataset,让它不仅存储所有图像的平铺列表,还记录每个目录对应的图像索引范围,这样后续Sampler能精准区分不同目录的图像。
import os import numpy as np from PIL import Image as io import torch from torch.utils.data import Dataset, Sampler class MainDataset(Dataset): def __init__(self, img_dir, use_folder_name=False): # 按目录分组的图像列表、平铺图像列表、每个目录的索引边界 self.grouped_images, self.gt_images, self.group_bounds = self._load_main_dataset(img_dir) self.dataset_len = len(self.gt_images) self.use_folder_name = use_folder_name def __len__(self): return self.dataset_len def __getitem__(self, idx): img_path = self.gt_images[idx] img_name = self._get_name(img_path) gt = self._load_img(img_path) if gt is None: return None gt = torch.from_numpy(gt).permute(2, 0, 1) return gt, img_name def _get_name(self, img_path): if self.use_folder_name: return img_path.split(os.sep)[-2] else: return img_path.split(os.sep)[-1].split('.')[0] def _load_main_dataset(self, img_dir): # 处理单个文件的特殊情况 if not os.path.isdir(img_dir): single_group = [img_dir] return [single_group], single_group, [(0, 1)] grouped_images = [] flat_images = [] group_bounds = [] start_idx = 0 # 遍历所有目录,按目录收集图像 for root, dirs, files in os.walk(img_dir): dir_images = [] for file in files: if is_valid_file(file): dir_images.append(os.path.join(root, file)) # 只保留包含有效图像的目录 if dir_images: dir_images.sort() grouped_images.append(dir_images) flat_images.extend(dir_images) # 记录当前目录的索引范围(左闭右开) end_idx = start_idx + len(dir_images) group_bounds.append((start_idx, end_idx)) start_idx = end_idx return grouped_images, flat_images, group_bounds def _load_img(self, img_path): try: gt_image = io.imread(img_path) gt_image_bd = getBitDepth(gt_image) gt_image = np.array(gt_image).astype(np.float32) / ((2 ** (gt_image_bd / 3)) - 1) return gt_image except Exception as e: print(f"加载图像失败 {img_path}: {e}") return None def is_valid_file(file_name: str): valid_image_extensions = ['.jpg', '.jpeg', '.png', '.gif', '.bmp', '.tiff', '.tif'] return any(file_name.lower().endswith(ext) for ext in valid_image_extensions) # 假设getBitDepth是你已有的函数,这里补充一个示例实现(如果没有的话) def getBitDepth(image): if image.dtype == np.uint8: return 24 elif image.dtype == np.uint16: return 48 else: return 24
2. 自定义BatchSampler:按目录生成完整批次
实现DirectoryBatchSampler,遍历每个目录的索引范围,只生成完整大小的批次,跳过不足批次的目录和剩余图像。
class DirectoryBatchSampler(Sampler): def __init__(self, group_bounds, batch_size): self.group_bounds = group_bounds self.batch_size = batch_size self.batches = self._generate_valid_batches() def _generate_valid_batches(self): batches = [] for start, end in self.group_bounds: total_imgs = end - start # 计算当前目录能生成的完整批次数量 num_valid_batches = total_imgs // self.batch_size if num_valid_batches == 0: continue # 不足一个批次,跳过该目录 # 生成每个批次的索引列表 for i in range(num_valid_batches): batch_start = start + i * self.batch_size batch_end = batch_start + self.batch_size batches.append(list(range(batch_start, batch_end))) return batches def __iter__(self): return iter(self.batches) def __len__(self): return len(self.batches)
3. 加载数据并验证逻辑
使用自定义的Dataset和Sampler创建DataLoader,即可实现需求中的批次生成逻辑:
# 替换为你的图像根目录路径 sdr_img_dir = "./your_image_root_dir" # 初始化数据集 sequence_data_store = MainDataset(img_dir=sdr_img_dir, use_folder_name=True) # 初始化自定义Sampler,设置批次大小为7 sampler = DirectoryBatchSampler(sequence_data_store.group_bounds, batch_size=7) # 创建DataLoader(注意不需要设置batch_size,因为Sampler已经定义了批次) sequence_loader = DataLoader(sequence_data_store, sampler=sampler, num_workers=0, pin_memory=False) # 验证批次内容 for batch_idx, (gt_batch, name_batch) in enumerate(sequence_loader): print(f"批次 {batch_idx} 大小: {len(gt_batch)}") # 打印该批次的目录名,验证是否来自同一目录 print(f"所属目录: {name_batch[0]}")
逻辑说明
- 目录A有15张图像:生成2个批次(0-6、7-13),剩余第14张图像被丢弃
- 目录B有10张图像:生成1个批次(0-6,对应目录内的前7张),剩余3张被丢弃
- 目录C有3张图像:因不足1个批次,直接被跳过
内容的提问来源于stack exchange,提问作者irgendwii
相关产品推荐
相关产品推荐

