如何用torch.unfold创建带重叠区域的图像分块?
实现带重叠区域的图像分块修改方案
要实现1024×1024分块且上下左右各100像素重叠,核心是调整unfold的步长参数,并修正padding逻辑,确保所有图像区域都能被覆盖。以下是修改后的代码及关键说明:
修改后的代码
import torch import torch.nn.functional as F def create_patches(image, patch_size=1024, overlap=100): # 计算分块步长:块大小减去重叠区域 stride = patch_size - overlap # 获取图像尺寸 (通道数, 高度, 宽度) _, h, w = image.shape # 计算高度方向需要的padding,确保最后一个块能完整覆盖图像区域 num_patches_h = (h + stride - 1) // stride # 向上取整计算块数 required_h = (num_patches_h - 1) * stride + patch_size padding_h = max(0, required_h - h) padding_top = padding_h // 2 padding_bottom = padding_h - padding_top # 计算宽度方向需要的padding num_patches_w = (w + stride - 1) // stride required_w = (num_patches_w - 1) * stride + patch_size padding_w = max(0, required_w - w) padding_left = padding_w // 2 padding_right = padding_w - padding_left # 对图像进行对称padding,避免边缘块仅单侧有padding image = F.pad(image, (padding_left, padding_right, padding_top, padding_bottom)) # 执行分块:通道维度保持不变,高度和宽度维度按步长stride滑动取块 patches = image.unfold(0, 3, 3) # 通道维度展开,保持每个块的通道数不变 patches = patches.unfold(1, patch_size, stride) # 高度方向分块,步长为stride patches = patches.unfold(2, patch_size, stride) # 宽度方向分块,步长为stride # 整理输出格式为 (块行数, 块列数, 通道数, 块高度, 块宽度),按需调整 patches = patches.permute(1, 2, 0, 3, 4).contiguous() return patches
关键修改点说明
- 步长调整:将
unfold的步长参数从patch_size改为patch_size - overlap,让相邻块之间自然产生100像素的重叠区域。 - Padding逻辑优化:原代码仅对右下侧补padding,现在改为对称padding(上下、左右均分padding量),避免边缘块的padding分布不均;同时通过计算所需总高度/宽度,确保最后一个块能完整覆盖原图像的边缘区域。
- 输出格式整理:新增
permute操作将分块结果调整为更直观的维度顺序,方便后续处理,若不需要可自行移除。
内容的提问来源于stack exchange,提问作者Below the Radar
相关产品推荐
相关产品推荐

