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

Keras自定义3×3卷积层训练时梯度异常问题求助(TensorFlow后端)

问题原因分析

你遇到的梯度为None的错误,核心原因是在call方法里使用了K.get_value(self.kernel)将张量转换为NumPy数组进行修改——这个操作会脱离TensorFlow的计算图追踪,导致后续的卷积操作无法被自动微分系统识别,自然就没有梯度信息了。

修正方案:保持操作在计算图内

我们需要用TensorFlow原生的张量操作来修改卷积核,而不是取出数组修改。这里可以通过创建一个掩码(mask),将卷积核的右下角元素置为0,所有操作都在计算图中完成,这样梯度就能正常传播了。

另外,你的compute_output_shape方法还有个小错误:第二个维度应该是input_shape[2]-2(对应特征图的宽度),而不是重复使用input_shape[1]-2。

以下是修正后的完整代码:

from keras.layers import Layer, Input, ZeroPadding2D
from keras.models import Model
import tensorflow as tf
from keras import backend as K

class CornerConv2D(Layer):
    def __init__(self, filters, **kwargs):
        self.filters = filters
        self.kernel_size = (3,3)
        super(CornerConv2D, self).__init__(**kwargs)
    
    def build(self, input_shape):
        shape = self.kernel_size + (input_shape[-1], self.filters)
        self.kernel = self.add_weight(name='kernel', shape=shape, initializer='glorot_uniform')
        super(CornerConv2D, self).build(input_shape)
    
    def call(self, layer_input):
        # 创建掩码:3x3的张量,除了右下角(2,2)位置为0,其余为1
        mask = tf.ones_like(self.kernel)
        # 生成右下角所有通道和滤波器的索引
        indices = [[2, 2, i, j] for i in range(self.kernel.shape[2]) for j in range(self.kernel.shape[3])]
        # 将右下角对应位置置为0
        mask = tf.tensor_scatter_nd_update(
            mask,
            indices=indices,
            updates=tf.zeros((self.kernel.shape[2]*self.kernel.shape[3],), dtype=self.kernel.dtype)
        )
        # 应用掩码到卷积核
        custom_kernel = self.kernel * mask
        # 执行卷积,配合外部的ZeroPadding2D,使用valid padding
        return K.conv2d(layer_input, custom_kernel, padding='valid')
    
    def compute_output_shape(self, input_shape):
        return (input_shape[0], input_shape[1]-2, input_shape[2]-2) + (self.filters,)

# 构建模型
inp = Input(shape=(64,64,1))
x = ZeroPadding2D(1)(inp)
x = CornerConv2D(30)(x)
x = ZeroPadding2D(1)(x)
x = CornerConv2D(30)(x)
x = ZeroPadding2D(1)(x)
x = CornerConv2D(1)(x)
m = Model(inputs=[inp], outputs=[x])
m.compile(loss='mse', optimizer="adam")
m.summary()
关键说明
  1. 保持计算图内操作:用tf.tensor_scatter_nd_update创建掩码并修改卷积核,所有操作都基于TensorFlow张量,没有脱离计算图,因此梯度可以正常被追踪。
  2. 卷积padding设置:因为你在自定义层外面已经添加了ZeroPadding2D(1),所以卷积时用padding='valid'就能得到正确的输出尺寸。
  3. 输出形状修正:将input_shape[1]-2改为input_shape[2]-2,确保宽高维度的计算正确。

现在你再运行m.fit训练代码,应该就能正常进行梯度更新了。

内容的提问来源于stack exchange,提问作者Moran Reznik

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 13:37:53