为何torch.nn.Conv2d处理多通道图像时会将图像分割为9块?
多通道卷积异常问题的解决方法
首先明确问题根源:你使用的nn.Conv2d默认是随机初始化权重,并非预设的常用滤波核(比如均值、高斯模糊)。单通道场景下可能刚好随机权重的视觉效果符合预期,但多通道下随机权重会导致三个通道的输出混乱,最终显示出异常的分块效果。
要在多通道场景实现正常卷积效果,你需要手动指定卷积核的权重,以常见的均值滤波为例,具体修改方案如下:
代码修改示例
import torch from torch import nn import cv2 import numpy as np # 读取图像并转换为PyTorch要求的NCHW格式 img = cv2.imread("image_game/eldenring 2022-12-14 19-29-50.png") cv2.imshow('input', img) # 转换为(batch_size, channels, height, width)格式 img_tensor = torch.tensor(img.transpose(2, 0, 1), dtype=torch.float32).unsqueeze(0) # 定义卷积层 c1 = nn.Conv2d(3, 3, kernel_size=(3, 3), padding=2, stride=1) # 手动设置均值滤波核,固定权重避免随机值干扰 with torch.no_grad(): # 生成3x3均值核,每个位置权重为1/9 kernel = torch.ones((3, 3)) / 9.0 # 扩展为卷积层权重要求的形状:(输出通道数, 输入通道数, 核高, 核宽) c1.weight.data = kernel.repeat(3, 3, 1, 1) # 偏置项设为0 c1.bias.data.zero_() # 执行卷积操作 output_tensor = c1(img_tensor) # 转换为cv2可显示的HWC格式,并处理数值范围 output = output_tensor.squeeze(0).transpose(1, 2, 0).detach().numpy() # 截断超出0-255的数值并转换为uint8格式,保证正常显示 output = np.clip(output, 0, 255).astype(np.uint8) cv2.imshow('output', output) cv2.waitKey(0) cv2.destroyAllWindows()
关键说明
- 权重初始化:多通道卷积的权重形状为
(out_channels, in_channels/groups, kernel_h, kernel_w),这里让每个输出通道对所有输入通道使用相同的均值核,保证色彩输出的一致性。 - 数值归一化:卷积后的输出可能超出0-255的图像显示范围,用
np.clip截断后转成uint8格式,才能让cv2正常渲染。 - 固定权重:用
torch.no_grad()包裹权重设置操作,避免后续如果涉及训练流程时权重被意外更新。
内容的提问来源于stack exchange,提问作者Karol Szymczak
相关产品推荐
相关产品推荐

