You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何理解Keras中layer.get_weights()返回的Conv2D层权重形状?

理解Keras Conv2D层layer.get_weights()的权重结构(以LeNet-5为例)

一、get_weights()返回的数组形状解析

Conv2D层调用get_weights()会返回两个数组:第一个是卷积核权重,第二个是偏置项(bias)。你提到的(5, 5, 1, 6)是卷积核权重的形状,各维度含义对应:

  • 5, 5:卷积核的高度(行)和宽度(列),对应LeNet-5的5x5卷积核规格
  • 1:输入特征图的通道数,你的输入是单通道灰度图(input_shape=(28,28,1)),所以这个维度固定为1
  • 6:输出的卷积核总数(filters参数),也就是LeNet-5第一个卷积层的6个独立卷积核

二、权重数组与卷积核的映射关系

你看到的嵌套数组结构,本质是NumPy多维数组的索引逻辑,对应关系非常直接:
对于形状为(kernel_height, kernel_width, input_channels, num_filters)的权重数组:

  • 最后一个维度num_filters对应不同的卷积核:索引0是第一个卷积核,1是第二个,以此类推到索引5的第六个卷积核
  • 前两个维度kernel_height, kernel_width对应卷积核内部的行和列:注意索引从0开始,比如[2,2]就对应卷积核的第三行第三列
  • 第三个维度input_channels对应输入通道:因为是单通道,这个维度始终取0

举两个具体操作的例子:

  • 提取第一个卷积核的所有权重:weights[..., 0](或weights[:, :, 0, 0]),得到的是一个(5,5)的数组,正好是一个完整的5x5卷积核
  • 访问第三个卷积核(索引为2)第三行第三列的权重:直接取weights[2, 2, 0, 2]

三、可视化卷积核的示例代码

用Matplotlib可以把6个卷积核直观可视化,对应权重结构:

import matplotlib.pyplot as plt

# 获取第一个Conv2D层的卷积核权重(取返回数组的第一个元素)
conv1_weights = model.layers[0].get_weights()[0]

# 批量可视化6个卷积核
plt.figure(figsize=(12, 8))
for idx in range(6):
    # 提取第idx个卷积核的权重,形状为(5,5)
    kernel = conv1_weights[:, :, 0, idx]
    plt.subplot(2, 3, idx+1)
    plt.imshow(kernel, cmap='gray')
    plt.title(f'卷积核 {idx+1}')
    plt.axis('off')
plt.tight_layout()
plt.show()

运行后会得到6张5x5的灰度图,每张图对应一个卷积核的权重分布,和weights[..., idx]的索引完全一一对应。

四、关于偏置项的补充

get_weights()返回的第二个数组是偏置项,形状为(6,),每个元素对应一个卷积核的偏置值,索引i对应第i+1个卷积核的偏置。


内容的提问来源于stack exchange,提问作者Syed Bilal Haider Bukhari

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.08 16:25:26