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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 17:05:55