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

修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 20:21:08