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

如何在Keras自定义损失函数中引入非训练集数据添加软约束

实现软约束的修正方案

核心错误说明

你原来的写法存在两个本质问题:

  • model.fit()是模型训练接口,不能用于在损失中获取其他样本的预测结果,且keras损失函数是基于静态计算图的张量运算,直接调用Python侧的模型预测/训练接口会打断计算图链路,无法正常反向传播梯度。
  • 损失函数默认仅能拿到当前批次的标注y_true和预测值y_pred,无法直接访问外部字典存储的supporters数据,需要把对应样本配对的supporter数据作为额外输入传入模型。

修正后实现代码

import numpy as np
import keras.backend as K
from keras.layers import Dense, Input
from keras.models import Model

# 1. 修正输入维度对齐:原来Input层是(20,)和features的(100,5)维度不匹配,这里统一为5维
input_dim = 5
output_dim = 3

# 构建基础特征提取网络(权重共享)
def build_shared_network(input_dim, output_dim):
    input_layer = Input((input_dim,))
    hidden_layer = Dense(16)(input_layer)
    output_layer = Dense(output_dim)(hidden_layer)
    return Model(inputs=input_layer, outputs=output_layer)

shared_net = build_shared_network(input_dim, output_dim)

# 2. 构建双输入模型:一个输入是原始训练特征,另一个是配对的supporter特征
train_input = Input((input_dim,), name="train_feature")
supporter_input = Input((input_dim,), name="supporter_feature")

# 共享权重计算两个输入的预测值
train_pred = shared_net(train_input)
supporter_pred = shared_net(supporter_input)

# 模型输出训练样本的预测值,用于计算MSE损失
model = Model(inputs=[train_input, supporter_input], outputs=train_pred)

# 3. 准备训练数据 + 配对的supporter数据
sample_count = 100
features = np.random.rand(sample_count, input_dim)
labels = np.random.rand(sample_count, output_dim)
# 把字典格式的supporters转成和训练样本一一对应的数组
supporters_arr = np.random.rand(sample_count, input_dim)

# 4. 自定义损失函数,加入软约束
def custom_loss(y_true, y_pred, supporter_pred):
    # 基础MSE损失
    mse = K.mean(K.square(y_true - y_pred), axis=-1)
    # 软约束:当前训练样本预测和配对supporter预测的最小绝对差
    soft_constraint = K.min(K.abs(y_pred - K.stop_gradient(supporter_pred)))
    # 可以调整lambda参数控制约束的权重
    lambda_constraint = 0.1
    return mse + lambda_constraint * soft_constraint

# 把supporter的预测值加入损失的计算参数
model.add_loss(custom_loss(labels, train_pred, supporter_pred))

# 5. 编译训练
model.compile(optimizer='sgd')
# 每批次传入1个训练样本+对应的supporter样本
model.fit(
    x=[features, supporters_arr],
    epochs=1,
    batch_size=1
)

关键逻辑说明

  • 用K.stop_gradient()固定supporter预测值对应的网络权重,满足你要求的「固定网络权重前提下计算差值」的需求,反向传播时不会更新supporter输入对应的梯度,仅优化约束项。
  • 采用共享权重的网络结构,保证训练样本和supporter样本用同一套权重做预测,符合需求逻辑。
  • 调整lambda_constraint的数值可以控制软约束在总损失中的权重,避免约束项覆盖原始MSE损失的优化目标。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 06:42:00