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

PyTorch自定义Dataset类实现训练时在线切分大图方法

PyTorch Dataset实时切图实现方案

核心逻辑是放弃预切图存盘的思路,在数据集初始化阶段提前算好所有子图的位置映射,读取原图时直接在内存中完成切片,不会产生额外磁盘占用,也不需要管理零散子图文件。

关键改动点

  • 新增切图配置参数:支持自定义子图尺寸、滑动步长(可配置重叠切图)
  • 预先生成子图索引映射表,训练时直接查表定位,避免运行时重复计算拖慢速度
  • 重写长度计算逻辑,返回所有子图的总数量而非原图数量
  • 修正原代码缺失依赖导入、缩进不规范的问题
  • 内置边缘适配逻辑:大图边缘不足一个子图尺寸时,自动调整最后一个切片的位置保证子图尺寸统一,不需要额外padding

修改后的完整代码

import os
import pandas as pd
import torch
from torch.utils.data import Dataset
from torchvision.io import read_image

class CustomImageDataset(Dataset):
    def __init__(self, annotations_file, img_dir, crop_size=512, step=None, transform=None, target_transform=None):
        """
        Args:
            annotations_file: 原图标签csv路径,格式和原有逻辑一致,第一列是文件名,第二列是标签
            img_dir: 原图存储目录
            crop_size: 切分的子图尺寸,传int表示正方形子图,传(h,w)元组表示自定义高宽
            step: 滑动窗口步长,默认等于crop_size即无重叠切图,小于crop_size时为重叠切图
            transform: 子图要应用的图像增强/预处理
            target_transform: 标签要应用的变换
        """
        self.img_labels = pd.read_csv(annotations_file)
        self.img_dir = img_dir
        self.transform = transform
        self.target_transform = target_transform
        
        # 处理切图参数
        if isinstance(crop_size, int):
            self.crop_h, self.crop_w = crop_size, crop_size
        else:
            self.crop_h, self.crop_w = crop_size
        self.step = step if step is not None else self.crop_h

        # 预生成所有子图的映射关系:每个元素格式为 (原图在csv中的索引, 子图左上角y坐标, 子图左上角x坐标)
        self.patch_map = []
        # 针对固定5000*5000尺寸的原图直接计算坐标,如果原图尺寸不固定,可提前把尺寸存在csv里逐图计算
        img_h, img_w = 5000, 5000
        # 生成y方向所有切片起点
        y_starts = list(range(0, img_h - self.crop_h + 1, self.step))
        # 处理边缘:如果最后一个切片没覆盖到图的底部,追加一个贴底的切片
        if y_starts[-1] + self.crop_h < img_h:
            y_starts.append(img_h - self.crop_h)
        # 生成x方向所有切片起点
        x_starts = list(range(0, img_w - self.crop_w + 1, self.step))
        if x_starts[-1] + self.crop_w < img_w:
            x_starts.append(img_w - self.crop_w)
        # 给每张原图绑定所有子图坐标
        for img_idx in range(len(self.img_labels)):
            for y in y_starts:
                for x in x_starts:
                    self.patch_map.append((img_idx, y, x))

    def __len__(self):
        # 返回总子图数量
        return len(self.patch_map)

    def __getitem__(self, idx):
        # 查表拿到当前子图对应的原图索引和切片坐标
        img_idx, y_start, x_start = self.patch_map[idx]
        # 读取对应原图
        img_path = os.path.join(self.img_dir, self.img_labels.iloc[img_idx, 0])
        image = read_image(img_path) # 读取后形状为 [C, H, W]
        label = self.img_labels.iloc[img_idx, 1]
        
        # 内存中直接切片取子图,不需要写盘
        # 注意CHW格式,通道维全取,后两个维度按坐标切
        patch = image[:, y_start:y_start+self.crop_h, x_start:x_start+self.crop_w]
        
        # 对子图应用变换
        if self.transform:
            patch = self.transform(patch)
        if self.target_transform:
            label = self.target_transform(label)
        return patch, label

使用示例

实例化数据集时指定子图尺寸即可,比如要切512*512无重叠子图:

# 无重叠切512*512子图
dataset = CustomImageDataset(
    annotations_file="train.csv",
    img_dir="./train_images",
    crop_size=512
)

# 如果要做重叠切图(比如分割任务减少边缘误差),指定步长小于子图尺寸即可,比如步长384,重叠128像素
dataset_overlap = CustomImageDataset(
    annotations_file="train.csv",
    img_dir="./train_images",
    crop_size=512,
    step=384
)

注意事项

  • 如果数据集里原图尺寸不固定,不要硬编码img_h和img_w,在__init__生成patch_map的环节逐张读取图片尺寸(或者提前把尺寸存在csv里)再计算切片坐标即可
  • 包含随机缩放、随机裁剪类的图像增强,建议放在切片操作之后执行,避免破坏预设的切图逻辑;归一化、颜色抖动类的增强不受顺序影响
  • 这个方案全程在内存中完成切片,Dataloader开多worker加载时不会有文件读写冲突,速度和预切图加载基本没有差异

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 02:45:35