CNN卷积层可视化疑问:滤波器维度转换与归一化方式探讨
CNN卷积层可视化相关问题解答
1. RGB卷积滤波器维度置换的正确性
你的做法完全正确。PyTorch中卷积层权重的默认维度是(out_channels, in_channels, height, width),针对输入为RGB图像的第一层卷积,in_channels=3对应RGB三通道。而Matplotlib的imshow函数要求RGB图像维度为(height, width, channels),因此通过permute(1, 2, 0)将filters[i]的维度从(in_ch, h, w)转换为(h, w, in_ch),完全匹配图像可视化的格式要求,能正确呈现RGB滤波器的样式。
2. 归一化方式的选择:全局vs单滤波器独立归一化
两种方式各有适用场景,取决于你想观察的核心信息:
- 全局归一化(当前代码的做法):用整个权重张量的最小/最大值统一归一化所有滤波器。这种方式会保留不同滤波器间的权重幅度差异——权重整体偏大的滤波器可视化后更亮,偏小的则更暗。但如果不同滤波器的权重范围差异悬殊,部分滤波器的细节可能被压缩,导致模式辨识度降低。
- 单滤波器独立归一化:对每个滤波器单独计算最小/最大值后映射到[0,1]。这种方式能最大化单个滤波器的对比度,让你清晰看到每个滤波器的权重分布模式,但会丢失滤波器间的权重幅度信息,无法直观对比响应强度的高低。
若要切换到独立归一化,可修改代码如下:
# 替换原归一化代码段 filters = [] for weight in weights: w_min, w_max = weight.min(), weight.max() normalized = (weight - w_min) / (w_max - w_min) filters.append(normalized) filters = torch.stack(filters)
原可视化代码
# Extract weights from the first layer weights = model.cn1.weight.detach().cpu() # Normalize weights to [0, 1] for visualization weights_min, weights_max = weights.min(), weights.max() filters = (weights - weights_min) / (weights_max - weights_min) fig, axes = plt.subplots(4, 8, figsize=(12, 6)) fig.suptitle('Filters of the First Convolutional Layer (cn1)', fontsize=16) for i in range(32): ax = axes[i // 8, i % 8] # Weights are (out_ch, in_ch, h, w). We permute to (h, w, in_ch) for plotting filter_img = filters[i].permute(1, 2, 0).numpy() ax.imshow(filter_img) ax.axis('off') plt.tight_layout(rect=[0, 0.03, 1, 0.95]) plt.show() def get_feature_maps(model, x): maps = [] # Block 1 x = F.relu(model.bn1(model.cn1(x))) x = F.max_pool2d(x, 2) maps.append(x.detach().cpu()) # Block 2 x = F.relu(model.bn2(model.cn2(x))) x = F.max_pool2d(x, 2) maps.append(x.detach().cpu()) # Block 3 x = F.relu(model.bn3(model.cn3(x))) x = F.max_pool2d(x, 2) maps.append(x.detach().cpu()) return maps test_img, test_label = next(iter(test_dl)) sample_input = test_img[0:1].to(device) # Shape (1, 3, 32, 32) feature_maps = get_feature_maps(model, sample_input) fig, axes = plt.subplots(3, 8, figsize=(15, 6)) for block_idx, fmap in enumerate(feature_maps): for i in range(8): ax = axes[block_idx, i] ax.imshow(fmap[0, i], cmap='viridis') ax.axis('off') if i == 0: ax.set_ylabel(f'Block {block_idx+1}', size='large') plt.suptitle('8 Channels of Feature Maps from each Conv Block', fontsize=16) plt.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者Sandesh Singh
相关产品推荐
相关产品推荐

