如何从带重叠的图像块还原图像(含超分辨率场景)
从重叠图像块还原原图像(含超分辨率场景)
你通过create_patches函数将形状为(c,h,w)的图像分割成带重叠的(i,j,c,h,w)形状图像块,现在需要反转该过程,尤其是超分辨率模型推理后(图像块形状变为(i,j,c,h*scale_factor,w*scale_factor))的还原,以下是具体实现方案:
核心思路
原分割过程分为padding补全图像和滑动窗口分块两步,还原时需要:
- 按原滑动窗口的逆过程,将每个图像块放回对应位置
- 对重叠区域做加权处理(避免拼接痕迹)
- 去掉原分割时添加的padding,还原到目标尺寸
具体实现代码
import torch def reconstruct_from_patches(patches, scale_factor, patch_size=224, overlap=24, org_image_shape=(3,669,1046)): org_c, org_h, org_w = org_image_shape # 计算基础参数 step = patch_size - overlap scaled_step = step * scale_factor scaled_patch_h = patch_size * scale_factor scaled_patch_w = patch_size * scale_factor i, j, _, _, _ = patches.shape # 计算补全后的超分辨率图像尺寸 scaled_padded_h = (i - 1) * scaled_step + scaled_patch_h scaled_padded_w = (j - 1) * scaled_step + scaled_patch_w # 初始化累加图像和权重掩码 output = torch.zeros((org_c, scaled_padded_h, scaled_padded_w), device=patches.device) weight_mask = torch.zeros((org_c, scaled_padded_h, scaled_padded_w), device=patches.device) # 生成线性衰减的权重(优化重叠区域拼接效果) def get_1d_weight(length, overlap_scaled): weight = torch.ones(length, device=patches.device) # 重叠区域线性衰减/递增,避免拼接痕迹 if overlap_scaled > 0: weight[:overlap_scaled] = torch.linspace(0, 1, overlap_scaled, device=patches.device) weight[-overlap_scaled:] = torch.linspace(1, 0, overlap_scaled, device=patches.device) return weight overlap_scaled = overlap * scale_factor h_weight = get_1d_weight(scaled_patch_h, overlap_scaled).unsqueeze(1) w_weight = get_1d_weight(scaled_patch_w, overlap_scaled).unsqueeze(0) patch_weight = h_weight * w_weight patch_weight = patch_weight.unsqueeze(0).repeat(org_c, 1, 1) # 扩展到通道维度 # 遍历所有patch,累加图像和权重 for i_idx in range(i): for j_idx in range(j): h_start = i_idx * scaled_step h_end = h_start + scaled_patch_h w_start = j_idx * scaled_step w_end = w_start + scaled_patch_w current_patch = patches[i_idx, j_idx] output[:, h_start:h_end, w_start:w_end] += current_patch * patch_weight weight_mask[:, h_start:h_end, w_start:w_end] += patch_weight # 处理权重掩码的零值(理论上不会出现) weight_mask[weight_mask == 0] = 1.0 # 归一化得到最终补全图像 output = output / weight_mask # 计算原分割时添加的padding尺寸,并裁剪还原到原图像超分辨率后的尺寸 padding_width = (step - (org_w - overlap) % step) % step padding_height = (step - (org_h - overlap) % step) % step scaled_pad_w = padding_width * scale_factor scaled_pad_h = padding_height * scale_factor # 裁剪掉右侧和底部的padding final_output = output[:, :org_h * scale_factor, :org_w * scale_factor] return final_output # 示例使用 # 假设超分辨率后的patches形状为torch.Size([4,6,3,896,896]) # org_image_shape是原图像的形状(3,669,1046) # restored_img = reconstruct_from_patches(sr_patches, scale_factor=4, org_image_shape=(3,669,1046))
关键细节说明
- 权重加权:使用线性衰减权重处理重叠区域,相比简单平均能大幅减少拼接痕迹,也可替换为余弦窗口等更平滑的权重函数
- padding处理:严格对应原分割函数的padding计算逻辑,确保裁剪后的尺寸与原图像超分辨率后的尺寸完全匹配
- 设备兼容:代码自动适配输入patches的设备(CPU/GPU),无需额外修改
内容的提问来源于stack exchange,提问作者Below the Radar
相关产品推荐
相关产品推荐

