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

在Keras Functional API LSTM中仅为自定义损失函数引入额外真值

解决方案

核心思路

  • 将额外真值y_train_v2作为辅助输入层加入模型,但不与任何训练层(LSTM/Dense)连接,仅用于损失函数计算
  • Keras在训练时会自动按批次对齐所有输入与标签,确保y_train和y_train_v2的批次完全匹配
  • 自定义损失函数通过包装器接收辅助输入层的张量,实现同时使用y_true、y_pred和y_true_v2计算损失

修改后的完整代码

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras.layers import Input, LSTM, Dense
from tensorflow.keras import backend as K

def custom_loss_wrapper(y_true_v2_input):
    def custom_loss(y_true, y_pred):
        # 按需结合原真值、预测值、额外真值计算损失
        # 示例:同时考虑两类真值与预测值的误差加权和
        loss_original = K.sqrt(K.mean(K.square(y_pred - y_true)))
        loss_v2 = K.sqrt(K.mean(K.square(y_pred - y_true_v2_input)))
        return 0.5 * loss_original + 0.5 * loss_v2  # 可自定义权重比例
    return custom_loss

def baseline_model():
    # 主输入:用于训练的X数据集
    main_input = Input(shape=(14, 1), name='main_input')
    # 辅助输入:仅用于损失计算的额外真值y_train_v2
    y_true_v2_input = Input(shape=(1,), name='y_true_v2_input')
    
    # 原模型结构保持不变,仅基于主输入构建
    x = LSTM(64, return_sequences=False)(main_input)
    output = Dense(1)(x)
    
    # 构建包含双输入、单输出的模型
    model = keras.Model(inputs=[main_input, y_true_v2_input], outputs=output)
    
    # 将辅助输入传入损失函数包装器
    model.compile(optimizer='adam', loss=custom_loss_wrapper(y_true_v2_input))
    return model

# 假设x_train/y_train/y_train_v2、x_test/y_test/y_test_v2已准备完成
model = baseline_model()
# 训练时输入为[主输入X, 额外真值y_train_v2],标签为原真值y_train
model.fit(
    [x_train, y_train_v2], 
    y_train, 
    epochs=10, 
    batch_size=4,
    validation_data=([x_test, y_test_v2], y_test)
)

关键细节说明

  • 辅助输入层的独立性:该输入层仅作为额外真值的传递载体,未连接任何可训练层,因此不会参与模型权重更新,完全满足"仅用于损失计算"的需求
  • 批次自动对齐:Keras的fit方法会自动拆分输入列表中的所有元素,保证同一批次内的x_train、y_train、y_train_v2严格对应,无需手动处理批次匹配问题
  • 损失函数灵活性:在custom_loss函数中,可根据业务逻辑自由组合三类张量(原真值、预测值、额外真值)计算最终损失,示例中的加权方式可按需调整

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 20:11:18