如何用PyTorch实时分割多分辨率图像为非重叠补丁喂给ResNet18
回答
1. PyTorch是否有现成的补丁分割Transform?
PyTorch官方torchvision.transforms模块中没有直接实现非重叠等大小图像补丁分割的现成Transform类。官方提供的RandomCrop、CenterCrop等工具仅能裁剪单块区域,无法一次性生成所有非重叠补丁,因此需要自定义Transform来实现需求。
2. 自定义补丁分割Transform及适配方案
结合你已有的ImageResizer类,我们可以实现Patch类完成分割,但需要注意:transforms.Compose默认期望每个Transform输出单个样本,而分割补丁会得到多个样本,直接放入Compose会与后续的ToTensor、DataLoader批量加载逻辑冲突,以下是两种可行方案:
方案一:实现Patch类,配合自定义Dataset处理多样本
首先完善ImageResizer(添加静态方法装饰器)并实现Patch类:
import numpy as np from PIL import Image import torch from torchvision import transforms from torch.utils.data import Dataset, DataLoader, ImageFolder class ImageResizer: """ 将图像调整为224的整数倍尺寸,方便后续分割为224x224的补丁 """ def __init__(self): pass @staticmethod def get_new_dimensions(width : int, height : int, patch_height : int = 224, patch_width : int = 224): """ 计算调整后的图像尺寸 参数: - width: 原图像宽度 - height: 原图像高度 - patch_height: 补丁高度 - patch_width: 补丁宽度 返回: - new_height: 调整后的高度 - new_width: 调整后的宽度 """ width_coef = int(np.round(width / patch_width)) height_coef = int(np.round(height / patch_height)) new_width = width_coef * patch_width new_height = height_coef * patch_height return new_width, new_height def __call__(self, image): width, height = image.size new_width, new_height = ImageResizer.get_new_dimensions(width, height) resized_image = image.resize((new_width, new_height)) return resized_image class Patch: """ 将图像分割为指定大小的非重叠补丁 """ def __init__(self, patch_size=(224, 224)): self.patch_h, self.patch_w = patch_size def __call__(self, image): # 将PIL图像转为numpy数组便于分割 img_np = np.array(image) h, w, c = img_np.shape # 计算垂直和水平方向的补丁数量 num_patches_h = h // self.patch_h num_patches_w = w // self.patch_w patches = [] # 遍历所有补丁区域 for i in range(num_patches_h): for j in range(num_patches_w): patch_np = img_np[i*self.patch_h : (i+1)*self.patch_h, j*self.patch_w : (j+1)*self.patch_w, :] # 转回PIL图像 patches.append(Image.fromarray(patch_np)) return patches
然后自定义Dataset来处理补丁分割后的多样本:
class PatchImageFolder(Dataset): def __init__(self, root, patch_size=(224,224)): self.base_dataset = ImageFolder(root) self.resizer = ImageResizer() self.patcher = Patch(patch_size) self.to_tensor = transforms.ToTensor() # 预生成所有补丁与标签的对应关系 self.patch_samples = [] for idx in range(len(self.base_dataset)): img, label = self.base_dataset[idx] resized_img = self.resizer(img) patches = self.patcher(resized_img) for patch in patches: self.patch_samples.append((patch, label)) def __len__(self): return len(self.patch_samples) def __getitem__(self, idx): patch, label = self.patch_samples[idx] return self.to_tensor(patch), label # 使用自定义Dataset加载数据 dataset = PatchImageFolder(root="<path>") batch_size = 32 dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
方案二:在Transform中直接输出张量补丁列表(需修改DataLoader的批量处理逻辑)
若不想自定义Dataset,可修改Patch类直接输出张量列表,但需要注意DataLoader默认会将列表堆叠为批量,此时需自定义collate_fn:
class PatchToTensor: def __init__(self, patch_size=(224,224)): self.patcher = Patch(patch_size) self.to_tensor = transforms.ToTensor() def __call__(self, image): patches = self.patcher(image) return [self.to_tensor(patch) for patch in patches] # 定义自定义collate_fn,将每个样本的补丁列表展开为批量 def patch_collate_fn(batch): # batch中的每个元素是(补丁列表, 标签) all_patches = [] all_labels = [] for patches, label in batch: all_patches.extend(patches) all_labels.extend([label]*len(patches)) return torch.stack(all_patches), torch.tensor(all_labels) # 构建transform transform = transforms.Compose([ ImageResizer(), PatchToTensor(patch_size=(224,224)), ]) dataset = ImageFolder(root="<path>", transform=transform) dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=patch_collate_fn)
关键优化点
- 你的
ImageResizer中get_new_dimensions方法已改为静态方法,避免了实例化调用的问题。 - 若需要保留图像原始比例,可修改
ImageResizer逻辑:先按最长边缩放至补丁大小的整数倍,再对短边进行补边(使用transforms.Pad),而非直接拉伸图像。
内容的提问来源于stack exchange,提问作者Yassine
相关产品推荐
相关产品推荐

