如何在TF2.0中创建带自定义梯度与可学习参数的Keras层?
问题修复方案
问题出在自定义梯度函数中,针对可学习参数s返回的梯度形状与参数本身形状不匹配:
- 可学习参数
scale的形状是[1,] - 当前返回的
dy_ds = x,形状是(100000,1),和参数形状维度不匹配,导致优化器更新时抛出形状错误。
解决方法是对dy_ds做归约操作(比如求和或平均),将其形状压缩为与scale一致的[1,]。因为每个样本对scale的梯度是对应输入x的值,总梯度需要是所有样本梯度的累加(对应MSE损失的梯度传播逻辑)。
修改后的代码如下:
# Method for calculation custom gradient @tf.custom_gradient def scaler(x, s): def grad(upstream): dy_dx = s # 对所有样本的x求和,将形状从(100000,1)转为(1,),和scale参数形状匹配 dy_ds = tf.reduce_sum(x * upstream, axis=0) return dy_dx, dy_ds return x * s, grad # Keras Layer with trainable parameter class TestLayer(tf.keras.layers.Layer): def build(self, input_shape): self.scale = self.add_weight("scale", shape=[1,], initializer=tf.keras.initializers.Constant(value=2.0), trainable=True) def call(self, inputs): return scaler(inputs, self.scale) # Creates Keras Model that uses the layer def Model(): x_in = tf.keras.layers.Input(shape=(1,)) x_out = TestLayer()(x_in) return tf.keras.Model(inputs=x_in, outputs=x_out, name="fp8_test") # Create toy dataset, want to learn `scale` such to satisfy 5 = 2 * scale (i.e, `scale` should learn ~2.5) def Dataset(): inps = tf.ones(shape=(10**5,1)) * 2 # 修正shape为(10**5,1),和输入层shape匹配 expected = tf.ones(shape=(10**5,1)) * 5 # 同样修正shape data_in = tf.data.Dataset.from_tensor_slices(inps) data_exp = tf.data.Dataset.from_tensor_slices(expected) dataset = tf.data.Dataset.zip((data_in, data_exp)).batch(32) # 添加batch,训练更高效 return dataset model = Model() model.summary() dataset = Dataset() # Use `MSE` loss and `SGD` optimizer model.compile( loss=tf.keras.losses.MSE, optimizer=tf.keras.optimizers.SGD(learning_rate=0.001), # 调整学习率,让收敛更稳定 ) model.fit(dataset, epochs=10)
额外说明:
- 数据集部分做了小修正:将
from_tensors改为from_tensor_slices并添加batch,避免一次性传入超大张量,同时让输入形状和模型输入层完全匹配。 - 调整了SGD的学习率,默认学习率可能过大导致训练不稳定。
训练后你会看到scale参数逐渐趋近于2.5,符合预期。
内容的提问来源于stack exchange,提问作者ohkneel
相关产品推荐
相关产品推荐

