TensorFlow 2.10中传入模型输入的自定义损失函数报错问题
问题解决:TensorFlow自定义损失函数同时接收输入、y_true、y_pred的实现问题
问题背景
使用TensorFlow 2.10实现自定义损失函数时,尝试将模型输入、y_true和y_pred同时传入损失函数,编写代码如下:
inputs = tf.keras.layers.Input(shape=X.shape[-1:], batch_size=16) dense1 = tf.keras.layers.Dense(8, activation='relu')(inputs) outputs = tf.keras.layers.Dense(1, activation='sigmoid')(dense1) model = tf.keras.Model(inputs=inputs, outputs=outputs) def custom_loss(i): def loss(y_true, y_pred): mask = tf.equal(y_true, tf.cast(tf.round(y_pred), 'float32')) selec = tf.where(mask, 1., 0.) x1 = (1 - tf.reduce_mean(selec)) mask = tf.reshape(selec, (-1,)) m = tf.boolean_mask(i, mask) x2 = tf.reduce_mean(m[:, -1]) return x1 * x2 return loss model.compile(loss=custom_loss(inputs), optimizer='adam', metrics=['accuracy', tf.keras.metrics.Precision(name='precision')]) h = model.fit(X_train, y_train, validation_data=(X_test, y_test), epochs=5, batch_size=16, verbose=2)
训练时出现第一个错误:
TypeError: You are passing KerasTensor(type_spec=TensorSpec(shape=(), dtype=tf.float32, name=None), name='Placeholder:0', description="created by layer 'tf.cast_17'"), an intermediate Keras symbolic input/output, to a TF API that does not allow registering custom dispatchers, such as `tf.cond`, `tf.function`, gradient tapes, or `tf.map_fn`. Keras Functional model construction only supports TF API calls that *do* support dispatching, such as `tf.math.add` or `tf.reshape`. Other APIs cannot be called directly on symbolic Kerasinputs/outputs. You can work around this limitation by putting the operation in a custom Keras layer `call` and calling that layer on this symbolic input/output.
尝试禁用即刻执行后,出现第二个错误:
ValueError: Variable <tf.Variable 'dense_14/kernel:0' shape=(13, 8) dtype=float32> has `None` for gradient. Please make sure that all of your ops have a gradient defined (i.e. are differentiable). Common ops without gradient: K.argmax, K.round, K.eval.
错误原因
- 符号张量API兼容问题:直接将
Input符号张量传入损失函数,并使用tf.boolean_mask这类不支持Keras符号张量的API,导致Functional模式下的符号运算报错。 - 不可导操作导致梯度丢失:
tf.round是不可导操作,禁用即刻执行后,反向传播时无法计算梯度,导致变量梯度为None。
解决方案
核心修改思路
- 通过模型多输出传递原始输入,避免直接传递
Input符号张量; - 替换不可导的
tf.round为可导的近似操作; - 用张量乘法替代
tf.boolean_mask,规避符号张量API限制。
修改后完整代码
import tensorflow as tf import numpy as np # 模拟训练/测试数据(替换为你的真实数据) X_train = np.random.rand(100, 13) y_train = np.random.randint(0, 2, size=(100, 1)) X_test = np.random.rand(20, 13) y_test = np.random.randint(0, 2, size=(20, 1)) # 1. 修改模型为多输出:同时输出预测值和原始输入 inputs = tf.keras.layers.Input(shape=X_train.shape[-1:], batch_size=16) dense1 = tf.keras.layers.Dense(8, activation='relu')(inputs) outputs = tf.keras.layers.Dense(1, activation='sigmoid')(dense1) # 将原始输入作为第二个输出,供损失函数调用 model = tf.keras.Model(inputs=inputs, outputs=[outputs, inputs]) # 2. 定义可导的自定义损失函数 def custom_loss(y_true, model_outputs): y_pred, inputs = model_outputs # 替换不可导的tf.round:用可导的阈值判断(或用平滑近似如tf.sigmoid((y_pred-0.5)*100)) pred_round = tf.cast(y_pred > 0.5, tf.float32) mask = tf.equal(y_true, pred_round) selec = tf.where(mask, 1., 0.) x1 = (1 - tf.reduce_mean(selec)) # 替换tf.boolean_mask:用张量乘法实现掩码,避免符号张量API限制 selec_expanded = tf.expand_dims(selec, axis=-1) # 筛选出预测错误的样本输入 masked_inputs = inputs * (1 - selec_expanded) # 计算错误样本最后一列的均值,避免除以0 non_zero_count = tf.reduce_sum(1 - selec) x2 = tf.cond( non_zero_count > 0, lambda: tf.reduce_sum(masked_inputs[:, -1]) / non_zero_count, lambda: tf.constant(0.0, dtype=tf.float32) ) return x1 * x2 # 3. 编译模型:指定多输出的损失和权重 model.compile( loss=[custom_loss, None], # 仅第一个输出用自定义损失,第二个输出无损失 optimizer='adam', metrics=['accuracy', tf.keras.metrics.Precision(name='precision')], loss_weights=[1.0, 0.0] # 仅计算第一个损失的权重 ) # 4. 训练模型:对应多输出传递标签和输入 h = model.fit( X_train, [y_train, X_train], # 第一个是真实标签,第二个对应模型的输入输出 validation_data=(X_test, [y_test, X_test]), epochs=5, batch_size=16, verbose=2 )
关键修改说明
- 多输出传递输入:将原始输入作为模型的第二个输出,让损失函数能合法获取输入数据,避免直接传递
Input符号张量的API兼容问题; - 可导近似替换:用
tf.cast(y_pred > 0.5, tf.float32)替代tf.round,保证梯度可计算;若需要更平滑的梯度,可改用tf.sigmoid((y_pred - 0.5)*100)近似阶跃函数; - 掩码逻辑重构:用张量乘法
inputs * (1 - selec_expanded)筛选错误样本,替代tf.boolean_mask,适配Keras符号张量的运算规则。
内容的提问来源于stack exchange,提问作者Tom Moritz
相关产品推荐
相关产品推荐

