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

Keras训练时如何为每个batch定制关联对应样本supporter向量的专属损失函数

实现方案

你没法在损失函数内部直接获取当前训练样本的索引,也不建议在损失中调用model.predict()(会破坏计算图导致梯度无法回传),最稳妥的实现方式是把每个样本对应的supporter向量作为额外输入传入模型,在计算图内部完成约束项的计算,完整实现代码如下:

1. 整理匹配对应输入数据

把每个样本对应的两个supporter单独提取,和原始特征一一对应:

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

# 原始数据
features = np.random.rand(100, 5)
labels = np.random.rand(100, 2)
holder = np.random.rand(200, 5) 

# 第i个样本对应holder[2*i]和holder[2*i+1],提取为两个独立的输入数组
supporter1 = holder[::2] # 形状 (100,5),对应每个样本的第一个supporter
supporter2 = holder[1::2] # 形状 (100,5),对应每个样本的第二个supporter

2. 搭建支持多输入的共享权重模型

复用同一套网络权重计算supporter的预测值,不需要调用predict:

# 定义共享层,保证特征、supporter都用当前训练的同一套权重计算输出
hidden_dense = Dense(16)
output_dense = Dense(2)

# 三个输入:原始特征、对应第一个supporter、对应第二个supporter
input_feature = Input((5,), name='feature_input')
input_s1 = Input((5,), name='s1_input')
input_s2 = Input((5,), name='s2_input')

# 计算原始特征的预测值
feat_hidden = hidden_dense(input_feature)
y_pred = output_dense(feat_hidden)

# 计算两个supporter的预测值,复用相同权重,不会额外产生训练参数
s1_hidden = hidden_dense(input_s1)
s1_pred = output_dense(s1_hidden)
s2_hidden = hidden_dense(input_s2)
s2_pred = output_dense(s2_hidden)

# 定义模型
model = Model(inputs=[input_feature, input_s1, input_s2], outputs=y_pred)

3. 自定义损失并训练

使用add_loss直接在计算图中构造损失逻辑:

# 定义损失计算逻辑
y_true = Input((2,), name='label_input')
# 基础MSE损失
mse = K.mean(K.square(y_true - y_pred), axis=-1)
# 约束项:当前样本预测值与对应两个supporter预测值的差的和
# 如果需要避免正负值抵消,可改为K.sum(K.abs(y_pred - s1_pred)) + K.sum(K.abs(y_pred - s2_pred))
new_constraint = K.sum(y_pred - s1_pred) + K.sum(y_pred - s2_pred)
total_loss = mse + new_constraint

model.add_loss(total_loss)

# 编译训练,每个样本自动匹配对应的supporter,不需要手动判断索引
model.compile(optimizer='sgd')
model.fit([features, supporter1, supporter2], labels, epochs=1, batch_size=1)

该方案所有计算都在计算图内部完成,梯度可以正常回传,后续如果要调整batch_size也可以直接复用,不需要修改核心逻辑。

内容的提问来源于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 10:15:01