基于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
相关产品推荐
相关产品推荐

