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

自定义Keras Dense层添加条件后无梯度问题求解

问题描述

我想在Keras的Dense层输出上做个简单的条件判断:如果输出≤0.001就设为1,否则设为0。于是自己写了个MyDense自定义层,在call方法里用tf.where实现了这个逻辑,但跑代码的时候直接报错了:

ValueError: No gradients provided for any variable, check your graph for ops that do not support gradients, between variables ["<tf.Variable 'scope0/rnn/while/lstm_cell/kernel:0' shape=(3, 512) dtype=float32>", "<tf.Variable 'scope0/rnn/while/lstm_cell/recurrent_kernel:0' shape=(128, 512) dtype=float32>", "<tf.Variable 'scope0/rnn/while/lstm_cell/bias:0' shape=(512,) dtype=float32>", "<tf.Variable 'scope0/my_dense/kernel:0' shape=(128, 1) dtype=float32>", "<tf.Variable 'scope0/my_dense/bias:0' shape=(1,) dtype=float32>"] and loss Tensor("Sum:0", shape=(), dtype=float32)

我的自定义层代码如下:

class MyDense(Layer):
    def __init__(self, units, activation=None, use_bias=True, kernel_initializer='glorot_uniform', bias_initializer='zeros', kernel_regularizer=None, bias_regularizer=None, activity_regularizer=None, kernel_constraint=None, bias_constraint=None, apply_cond = False, **kwargs):
        if 'input_shape' not in kwargs and 'input_dim' in kwargs:
            kwargs['input_shape'] = (kwargs.pop('input_dim'),)
        super(MyDense, self).__init__(
            activity_regularizer=regularizers.get(activity_regularizer),
            **kwargs)
        self.units = int(units)
        self.activation = activations.get(activation)
        self.use_bias = use_bias
        self.kernel_initializer = initializers.get(kernel_initializer)
        self.bias_initializer = initializers.get(bias_initializer)
        self.kernel_regularizer = regularizers.get(kernel_regularizer)
        self.bias_regularizer = regularizers.get(bias_regularizer)
        self.kernel_constraint = constraints.get(kernel_constraint)
        self.bias_constraint = constraints.get(bias_constraint)
        self.apply_cond = apply_cond
        self.supports_masking = True
        self.input_spec = InputSpec(min_ndim=2)

    def build(self, input_shape):
        input_shape = tensor_shape.TensorShape(input_shape)
        if tensor_shape.dimension_value(input_shape[-1]) is None:
            raise ValueError('The last dimension of the inputs to `Dense` '
                             'should be defined. Found `None`.')
        last_dim = tensor_shape.dimension_value(input_shape[-1])
        self.input_spec = InputSpec(min_ndim=2, axes={-1: last_dim})
        self.kernel = self.add_weight(
            'kernel',
            shape=[last_dim, self.units],
            initializer=self.kernel_initializer,
            regularizer=self.kernel_regularizer,
            constraint=self.kernel_constraint,
            dtype=self.dtype,
            trainable=True)
        if self.use_bias:
            self.bias = self.add_weight(
                'bias',
                shape=[self.units,],
                initializer=self.bias_initializer,
                regularizer=self.bias_regularizer,
                constraint=self.bias_constraint,
                dtype=self.dtype,
                trainable=True)
        else:
            self.bias = None
        self.built = True

    def call(self, inputs):
        # print('in start of call apply_cond is: ', self.apply_cond)
        inputs = ops.convert_to_tensor(inputs)
        rank = common_shapes.rank(inputs)
        if rank > 2:
            # Broadcasting is required for the inputs.
            outputs = standard_ops.tensordot(inputs, self.kernel, [[rank - 1], [0]])
            # Reshape the output back to the original ndim of the input.
            if not context.executing_eagerly():
                shape = inputs.get_shape().as_list()
                output_shape = shape[:-1] + [self.units]
                outputs.set_shape(output_shape)
        else:
            outputs = gen_math_ops.mat_mul(inputs, self.kernel)
        if self.use_bias:
            outputs = nn.bias_add(outputs, self.bias)
        if self.activation is not None:
            outputs = self.activation(outputs)  # pylint: disable=not-callable
        if self.apply_cond:
            cond = tf.less_equal(outputs, tf.constant(0.00001), name='mycondition')
            return tf.where(cond, tf.ones_like(outputs), tf.zeros_like(outputs), name='mywhere')
        return outputs

    def compute_output_shape(self, input_shape):
        input_shape = tensor_shape.TensorShape(input_shape)
        input_shape = input_shape.with_rank_at_least(2)
        if tensor_shape.dimension_value(input_shape[-1]) is None:
            raise ValueError(
                'The innermost dimension of input_shape must be defined, but saw: %s'
                % input_shape)
        return input_shape[:-1].concatenate(self.units)
问题原因与修复方法

为啥会报错?

这个问题的核心是:你用tf.where实现的硬阈值判断是个不可导的操作。当你把输出直接切成0或1时,这个函数在阈值点的梯度根本不存在——反向传播的时候,模型没法计算权重的更新梯度,自然就会抛出"No gradients provided"的错误。

怎么修复?

我们需要把这个不可导的硬判断换成可导的近似函数,让梯度能正常在网络中传播。这里有两种常用的方案:

方案1:用陡峭的Sigmoid函数模拟硬阈值

Sigmoid函数本身是处处可导的,我们可以通过放大输入让它变得非常陡峭,效果就接近硬判断了。修改你的call方法里的条件判断部分:

def call(self, inputs):
    # 保留原来的计算逻辑(矩阵乘法、偏置、激活函数等)...
    if self.apply_cond:
        # 我们的需求是:outputs ≤0.001 → 1,否则0
        # 先把条件转换为(0.001 - outputs),这样符合条件时这个值≥0,sigmoid后接近1
        scale = 10000.0  # 缩放系数越大,函数越陡峭,越接近硬判断
        approx_output = tf.sigmoid(scale * (0.001 - outputs))
        return approx_output
    return outputs

方案2:训练用近似,推理用硬判断

如果你的场景要求推理时必须得到严格的0/1值,可以在call方法里区分训练和推理模式:

def call(self, inputs):
    # 保留原来的计算逻辑...
    if self.apply_cond:
        if tf.keras.backend.learning_phase():
            # 训练阶段:用可导的Sigmoid近似,保证梯度传播
            scale = 10000.0
            approx_output = tf.sigmoid(scale * (0.001 - outputs))
            return approx_output
        else:
            # 推理阶段:用原来的硬判断,得到严格的0/1
            cond = tf.less_equal(outputs, tf.constant(0.001), name='mycondition')
            return tf.where(cond, tf.ones_like(outputs), tf.zeros_like(outputs), name='mywhere')
    return outputs

小提示

  • 别用ReLU这类函数替代,因为ReLU在0点的梯度也是0,还是会导致梯度消失的问题,而Sigmoid的近似在整个区间都有非零梯度,能保证反向传播正常进行。
  • 缩放系数可以按需调整:系数越大,函数越接近硬阈值,但也要注意数值稳定性,避免出现梯度爆炸的情况。

内容的提问来源于stack exchange,提问作者Atr Cheema

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:37:11