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

基于U-net的深度学习数据集预处理:生成64x64步长32图像补丁

解决U-Net训练中512×512图像的64×64补丁分割问题

嘿,我懂你在做U-Net训练时碰到的数据集预处理难题了——要把512×512的图像切成步长32的64×64补丁对吧?这事儿用NumPy或者结合PyTorch的数据集类就能轻松搞定,我给你梳理个完整的解决方案,训练图、标签图、测试图都能通用:

核心思路

  • 先算清楚每个方向能生成多少补丁:512的边长,补丁大小64,步长32,那每个维度的补丁数是 (512 - 64) // 32 + 1 = 15,也就是每张图能切出15×15=225个补丁
  • 遍历每个补丁的起始坐标,直接对图像数组切片提取就行

基础NumPy实现(适用于单张/批量图像)

这个函数可以处理单通道(比如标签图)或多通道(比如RGB训练图)的图像,直接返回所有补丁的数组:

import numpy as np

def extract_patches(image, patch_size=64, stride=32):
    # 获取图像的高和宽,兼容单通道(H,W)或多通道(H,W,C)格式
    h, w = image.shape[:2]
    # 计算高、宽方向的补丁数量
    num_h = (h - patch_size) // stride + 1
    num_w = (w - patch_size) // stride + 1
    patches = []
    
    for i in range(num_h):
        for j in range(num_w):
            # 计算当前补丁的起始/结束坐标
            start_h = i * stride
            end_h = start_h + patch_size
            start_w = j * stride
            end_w = start_w + patch_size
            # 切片提取补丁
            patch = image[start_h:end_h, start_w:end_w]
            patches.append(patch)
    
    # 转换为numpy数组,形状为(补丁总数, patch_size, patch_size, 通道数)(多通道时)
    return np.array(patches)

# 示例使用
# 模拟512×512的单通道训练图像和标签
train_image = np.random.rand(512, 512)
label_image = np.random.randint(0, 2, size=(512, 512))

train_patches = extract_patches(train_image)
label_patches = extract_patches(label_image)

print(f"单张图生成的补丁数量:{train_patches.shape[0]}")  # 输出225
print(f"每个补丁的尺寸:{train_patches.shape[1:]}")       # 输出(64,64)

结合PyTorch的数据集类(适合训练时动态加载)

如果你的数据集很大,直接预生成所有补丁会占用太多内存,可以用这个动态加载的数据集类,训练时按需提取补丁:

from torch.utils.data import Dataset

class PatchDataset(Dataset):
    def __init__(self, images, labels, patch_size=64, stride=32):
        self.images = images
        self.labels = labels
        self.patch_size = patch_size
        self.stride = stride
        # 预存所有补丁的索引信息(对应原图下标+补丁坐标)
        self.patch_info = []
        for img_idx in range(len(images)):
            h, w = images[img_idx].shape[:2]
            num_h = (h - patch_size) // stride + 1
            num_w = (w - patch_size) // stride + 1
            for i in range(num_h):
                for j in range(num_w):
                    self.patch_info.append((img_idx, i, j))
    
    def __len__(self):
        return len(self.patch_info)
    
    def __getitem__(self, idx):
        img_idx, i, j = self.patch_info[idx]
        # 计算补丁坐标
        start_h = i * self.stride
        end_h = start_h + self.patch_size
        start_w = j * self.stride
        end_w = start_w + self.patch_size
        # 提取补丁
        img_patch = self.images[img_idx][start_h:end_h, start_w:end_w]
        label_patch = self.labels[img_idx][start_h:end_h, start_w:end_w]
        # 转换为PyTorch要求的张量格式:(通道数, 高, 宽)
        if len(img_patch.shape) == 3:
            img_patch = np.transpose(img_patch, (2, 0, 1))
        else:
            img_patch = img_patch[np.newaxis, :, :]
        label_patch = label_patch[np.newaxis, :, :]
        return img_patch.astype(np.float32), label_patch.astype(np.long)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:36:09