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

基于PyTorch实现自定义DataLoader:加载带可变窗口的图像与掩码并追踪目录

PyTorch自定义Dataset实现(适配可变尺寸Patch数据集)

下面是针对你描述的数据集结构,实现的自定义PyTorch Dataset类,支持加载可变尺寸的图像Patch、对应掩码,并追踪所属的主文件夹索引:

代码实现

import os
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader

class VariablePatchDataset(Dataset):
    def __init__(self, root_dir, image_size=None):
        """
        Args:
            root_dir (str): 数据集根目录(即All_data的路径)
            image_size (int, optional): 原始图像的边长(正方形)。如果不传入,会尝试从starts.npy推导最后一个patch的尺寸
        """
        self.root_dir = root_dir
        self.data_pairs = []
        
        # 遍历Numpy_dat下的所有子文件夹
        dat_root = os.path.join(root_dir, "Numpy_dat")
        mask_root = os.path.join(root_dir, "mask_numpy")
        dat_folders = sorted([f for f in os.listdir(dat_root) if f.startswith("dat_")])
        
        for dat_folder in dat_folders:
            # 提取主文件夹索引z
            z = int(dat_folder.split("_")[-1])
            dat_path = os.path.join(dat_root, dat_folder)
            mask_path = os.path.join(mask_root, f"mask_{z}")
            
            # 读取起始点数组
            starts = np.load(os.path.join(dat_path, "starts.npy"))
            # 计算每个patch的尺寸
            if image_size is not None:
                # 用传入的图像边长计算最后一个patch的尺寸
                patch_sizes = np.diff(starts, append=image_size)
            else:
                # 如果没有传入image_size,假设starts的最后一个元素+最后一个patch尺寸等于图像边长
                # 这里需要根据你的实际数据调整,如果starts包含结束点,可直接用diff
                patch_sizes = np.diff(starts)
                # 注意:如果starts只存起始点,这里需要补充最后一个patch的尺寸逻辑
                # 比如可以从第一个patch文件的形状推导,但不同patch尺寸不同,所以建议传入image_size
            
            # 遍历每个patch文件
            for idx in range(len(starts)):
                patch_file = os.path.join(dat_path, f"dat_{z}_{idx}.npy")
                mask_file = os.path.join(mask_path, f"mask_{z}_{idx}.npy")
                self.data_pairs.append({
                    "patch_path": patch_file,
                    "mask_path": mask_file,
                    "patch_size": patch_sizes[idx],
                    "folder_idx": z
                })
    
    def __len__(self):
        return len(self.data_pairs)
    
    def __getitem__(self, idx):
        item = self.data_pairs[idx]
        
        # 加载图像patch和掩码
        patch = np.load(item["patch_path"])
        mask = np.load(item["mask_path"])
        
        # 调整维度:如果是单通道,添加通道维度(比如从(H,W)转为(1,H,W))
        if len(patch.shape) == 2:
            patch = np.expand_dims(patch, axis=0)
        if len(mask.shape) == 2:
            mask = np.expand_dims(mask, axis=0)
        
        # 转换为PyTorch tensor
        patch_tensor = torch.tensor(patch, dtype=torch.float32)
        mask_tensor = torch.tensor(mask, dtype=torch.long)  # 分割任务用long,分类可改float32
        
        return {
            "patch": patch_tensor,
            "mask": mask_tensor,
            "patch_size": item["patch_size"],
            "folder_idx": item["folder_idx"]
        }

使用示例

1. 基础加载(单样本)

# 初始化数据集
dataset = VariablePatchDataset(root_dir="./All_data", image_size=254)
# 获取单个样本
sample = dataset[0]
print(f"Patch shape: {sample['patch'].shape}")
print(f"Mask shape: {sample['mask'].shape}")
print(f"Patch size: {sample['patch_size']}")
print(f"所属文件夹索引: {sample['folder_idx']}")

2. 结合DataLoader加载

由于patch尺寸可变,默认的DataLoader无法直接批量加载(会因形状不匹配报错),如果需要批量处理,需要自定义collate_fn进行padding:

def custom_collate_fn(batch):
    # 找到当前batch中最大的patch尺寸
    max_size = max(item["patch_size"] for item in batch)
    padded_patches = []
    padded_masks = []
    sizes = []
    folder_idxs = []
    
    for item in batch:
        patch = item["patch"]
        mask = item["mask"]
        # 计算需要padding的尺寸
        pad_h = max_size - patch.shape[1]
        pad_w = max_size - patch.shape[2]
        # 进行padding(默认填充0,可根据需求调整)
        padded_patch = torch.nn.functional.pad(patch, (0, pad_w, 0, pad_h))
        padded_mask = torch.nn.functional.pad(mask, (0, pad_w, 0, pad_h))
        
        padded_patches.append(padded_patch)
        padded_masks.append(padded_mask)
        sizes.append(item["patch_size"])
        folder_idxs.append(item["folder_idx"])
    
    return {
        "patches": torch.stack(padded_patches),
        "masks": torch.stack(padded_masks),
        "patch_sizes": torch.tensor(sizes),
        "folder_idxs": torch.tensor(folder_idxs)
    }

# 初始化DataLoader
dataloader = DataLoader(dataset, batch_size=4, shuffle=True, collate_fn=custom_collate_fn)

# 遍历DataLoader
for batch in dataloader:
    print(f"Batch patches shape: {batch['patches'].shape}")
    print(f"Batch masks shape: {batch['masks'].shape}")
    print(f"Batch patch sizes: {batch['patch_sizes']}")
    print(f"Batch folder indexes: {batch['folder_idxs']}")
    break

关键说明

  • 文件夹索引追踪:通过解析dat_z文件夹名称中的数字,获取每个patch所属的主文件夹索引,并在返回结果中提供。
  • 可变尺寸处理:通过starts.npy计算每个patch的尺寸,支持不同行列的patch尺寸差异。
  • 自定义collate_fn:针对可变尺寸的patch,实现批量加载时的padding逻辑,确保batch内的tensor形状一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 23:22:05