如何在Keras中为CNN回归模型自定义带时序惩罚的损失函数?
问题描述
我正在做一个基于视频图像的CNN回归模型,用来预测异常事件发生前的剩余次数。标签会随着接近异常事件递减(比如异常事件对应标签0,前一帧是1,再往前是2),所以预测值也得是递减的。我想在损失函数里加个惩罚项:要是当前预测值比前一次的大,就罚它,而且得用Keras实现,不能用tf.GradientTape()。
现有模型架构代码:
initial_model = tf.keras.applications.VGG16(weights='imagenet', include_top=False) initial_model.trainable = False func_model_p = keras.Sequential() inputs = keras.Input(shape=(700, 100, 3)) func_model_p.add(inputs) func_model_p.add(keras.layers.GlobalAveragePooling2D()) func_model_p.add(keras.layers.Dense(1, activation="linear")) # 准备指标 train_mse_metric_p = keras.metrics.MeanSquaredError() val_mse_metric_p = keras.metrics.MeanSquaredError()
我现在用tf.GradientTape()写出了逻辑,想改成Keras原生实现,代码如下:
for e in range(epochs): for i in range(len(y_train_d_noshuffle)): with tf.GradientTape() as tape: image = np.expand_dims(X_train_d_noshuffle[i], axis=0) y_hat = func_model_p(image, training=True) previous_value = previous_list[-1] violation_term = tf.constant(max( (y_hat - previous_value) , 0), dtype=tf.float32) try: if y_train_d_noshuffle[i] < y_train_d_noshuffle[i+1]: violation_term = 0 except: pass y_train_c = tf.constant(y_train_d_noshuffle[i]) mse = tf.keras.losses.MeanSquaredError()(y_train_c, y_hat) # 计算损失 loss_value = mse + violation_term previous_list.append(y_hat) train_mse_sum = train_mse_sum + mse train_penalty_sum = train_penalty_sum + violation_term grads = tape.gradient(loss_value, func_model_p.trainable_weights) optimizer.apply_gradients(zip(grads, func_model_p.trainable_weights)) # 更新训练指标 train_mse_metric_p.update_state(y_train_d_noshuffle[i], y_hat) # 在每个epoch结束时显示指标 train_mse = train_mse_metric_p.result().numpy() print(f'epochs: {e}, train_mse_sum: {train_mse_sum}, train_penalty_sum: {float(train_penalty_sum)}')
Keras原生实现方案
要实现这个带序列惩罚的逻辑,关键是要跟踪前一个样本的预测值,同时贴合Keras的训练流程。下面是具体实现步骤:
1. 补全模型结构
先提一句:你原模型里漏了VGG16特征提取层的接入,直接把输入接全局平均池化层是没法处理图像的,我下面的代码里补上了这部分。
2. 完整训练代码
import tensorflow as tf from tensorflow import keras import numpy as np # 初始化完整模型(补上遗漏的VGG16层) initial_model = tf.keras.applications.VGG16(weights='imagenet', include_top=False) initial_model.trainable = False func_model_p = keras.Sequential([ keras.Input(shape=(700, 100, 3)), initial_model, keras.layers.GlobalAveragePooling2D(), keras.layers.Dense(1, activation="linear") ]) # 准备指标 train_mse_metric = keras.metrics.MeanSquaredError() train_penalty_metric = keras.metrics.Mean() # 初始化优化器 optimizer = keras.optimizers.Adam() # 存前一次预测值的状态变量,初始值设很大,确保第一个样本不触发惩罚 prev_pred = tf.Variable([1e10], dtype=tf.float32) epochs = 10 # 自行替换你的epoch数 for e in range(epochs): # 每个epoch开头重置状态和指标 prev_pred.assign([1e10]) train_mse_metric.reset_state() train_penalty_metric.reset_state() train_mse_sum = 0.0 train_penalty_sum = 0.0 for i in range(len(y_train_d_noshuffle)): # 取当前样本和标签 x = np.expand_dims(X_train_d_noshuffle[i], axis=0) y_true = np.expand_dims(y_train_d_noshuffle[i], axis=0) violation_term = 0.0 # 计算预测值和惩罚项 with tf.GradientTape() as tape: y_pred = func_model_p(x, training=True) # 只有当前标签不小于下一个标签时,才计算惩罚 if i < len(y_train_d_noshuffle) - 1: if y_train_d_noshuffle[i] >= y_train_d_noshuffle[i+1]: violation = tf.maximum(y_pred - prev_pred, 0.0) violation_term = tf.reduce_mean(violation) else: # 最后一个样本,正常算惩罚 violation = tf.maximum(y_pred - prev_pred, 0.0) violation_term = tf.reduce_mean(violation) # 计算MSE损失和总损失 mse = tf.reduce_mean(tf.keras.losses.MSE(y_true, y_pred)) total_loss = mse + violation_term # 更新模型权重(Keras原生优化器调用) grads = tape.gradient(total_loss, func_model_p.trainable_weights) optimizer.apply_gradients(zip(grads, func_model_p.trainable_weights)) # 更新状态和指标 prev_pred.assign(y_pred) train_mse_metric.update_state(y_true, y_pred) train_penalty_metric.update_state(violation_term) train_mse_sum += mse.numpy() train_penalty_sum += violation_term.numpy() # 打印epoch结果 print(f"epoch: {e+1}, train_mse_sum: {train_mse_sum:.4f}, train_penalty_sum: {train_penalty_sum:.4f}, train_mse: {train_mse_metric.result().numpy():.4f}")
关键说明
- 补上了原模型中缺失的VGG16特征提取层,这是处理图像输入的核心,不然模型根本没法从图像里提取特征。
- 用
tf.Variable跟踪前一次预测值,每个epoch重置,避免跨epoch的状态污染。 - 保留了你原逻辑里“当当前标签小于下一个标签时不惩罚”的规则,只在标签递减的序列段施加惩罚。
- 用Keras的优化器和指标API,替代了手动的梯度计算循环,更符合Keras的使用规范。
内容的提问来源于stack exchange,提问作者Minsung Kang
相关产品推荐
相关产品推荐

