卷积神经网络中间层特征图的可视化

python 复制代码
import torch
import torch.nn as nn
import matplotlib.pyplot as plt
from torchvision import transforms
from PIL import Image


# 定义卷积神经网络模型
class SimpleCNN(nn.Module):
    def __init__(self):
        super(SimpleCNN, self).__init__()
        self.conv = nn.Conv2d(3, 8, kernel_size=3, stride=1, padding=1)
        self.bn = nn.BatchNorm2d(8)
        self.relu = nn.ReLU()
        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)

    def forward(self, x):
        x = self.conv(x)
        # x = self.bn(x)
        # x = self.relu(x)
        # x = self.pool(x)
        return x


if __name__ == '__main__':
    # 设置 CPU 张量的随机数种子
    torch.manual_seed(42)

    # 创建模型实例
    model = SimpleCNN()

    # 加载并预处理图片
    img_path = r'E:\photo\123.jpg'
    img = Image.open(img_path).convert('RGB')  # 读取的默认格式为 RGB,这里可去掉 convert()
    preprocess = transforms.Compose([transforms.Resize((960, 960)),
                                     transforms.ToTensor(),
                                     transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])
    img_tensor = preprocess(img).unsqueeze(0)  # (1, C, H, W)

    # 不计算梯度,进行一次前向传播
    with torch.no_grad():
        output = model(img_tensor)

    # 模型输出的图片大小
    print("Output size after conv layer:", output.size())

    # 可视化原始图片
    plt.imshow(img)
    plt.title("Original Image")
    plt.axis('off')
    plt.show()

    # 可视化卷积层后的图片
    for i in range(output.size()[1]):
        plt.subplot(output.size()[1]//4, 4, i+1)
        plt.imshow(output[0, i, :, :].cpu().detach().numpy())
        plt.axis('off')
    plt.tight_layout()
    plt.subplots_adjust(hspace=0.05)
    plt.suptitle('After Conv2d')
    plt.show()

原图大小为:(872, 1280, 3)

原图如下所示:

图片经过卷积层后,得到的特征图大小为:torch.Size(1, 8, 960, 960)

图片经过卷积层后,得到的特征图如下所示:

图片经过 BN 层后,得到的特征图大小为:torch.Size(1, 8, 960, 960)

图片经过 BN 层后,得到的特征图如下所示:

图片经过 ReLU 层后,得到的特征图大小为:torch.Size(1, 8, 960, 960)

图片经过 ReLU 层后,得到的特征图如下所示:

图片经过最大池化层后,得到的特征图大小为:torch.Size(1, 8, 480, 480)

图片经过最大池化层后,得到的特征图如下所示:

PIL 中的 Image.open(img_path) 读取的图片维度为 (W, H, C),读取的图片模式默认为 RGB;

OpenCV 中的 cv2.imread(img_path) 读取的图像维度为 (H, W, C),读取的图片模式默认为 BGR;

Image 图像数据转换为 np.ndarray 时,格式会从 (W, H, C) 转换为 (H, W, C)。

transforms.Resize((960, 960):旨在改变图像的大小,默认使用双线性插值(Bilinear);还支持最近邻插值(Nearest)、双三次插值(Bicubic)。

transforms.ToTensor():将 PIL 图像或 Numpy 数组中的整数像素值转换为 torch.FloatTensor 类型的浮点数;如果 PIL Image 属于 (L, LA, P, I, F, RGB, YCbCr, RGBA, CMYK, 1) 中的一种图像类型,或者 numpy.ndarray 的数据类型是 np.uint8,则将像素值从 0, 255 归一化到 0.0, 1.0,这是通过将每个像素值除以 255 来实现的;将 H, W, C 的图像格式转换为 C, H, W 的 tensor 格式。

相关推荐
大模型momo1 小时前
Spring AI 实战:多 Agent 协作实战 —— 分工拆解复杂旅游行程任务
人工智能·spring·ai·agent·旅游
小程故事多_801 小时前
从A2C、TRPO、PPO到GRPO,强化学习策略梯度算法完整演进与大模型落地实战解析
人工智能·算法
冬奇Lab2 小时前
开源项目第176期:Better Harness — 不审查 diff,审查工作流本身,给 AI 编程 Agent 的五维评估框架
人工智能·开源·agent
冬奇Lab2 小时前
代码库知识库系列(07):混合检索 BM25 + 向量——Q8 还是失败,而且总分退步了
人工智能
ajassi20002 小时前
AI语音智能体开发日记(十一)为智能设备“声”临其境——详解音频资源自动化生成流程
人工智能·ai·ai编程
2601_949499942 小时前
400G组网低功耗优选!芯瑞科技400G VR4 QSFP112光模块赋能智算中心高速互联
大数据·人工智能·科技
GoAI3 小时前
# AI Agent 记忆框架横向对比报告总结
人工智能·大模型·llm·多模态
AI人工智能+3 小时前
一种基于深度学习技术的高精度医疗机构执业许可证识别系统,构建了一套基于深度神经网络的端到端智能识别系统,为医疗行业提
深度学习·ocr·医疗机构执业许可证识别
硅谷秋水3 小时前
EgoSteer:一种基于第一人称视角视频、实现可控灵巧操作的全栈系统
深度学习·机器学习·语言模型·机器人·音视频
李昊哲小课3 小时前
fastapi sse websocket 奶茶店实时订单看板
人工智能·python·websocket·网络协议·fastapi·sse