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

Keras拟合多输入多输出自编码器报错及EarlyStopping疑问

报错原因与cardinality含义
  • 报错核心原因:model.fit()的参数传参错误。自编码器属于重构类模型,训练时的预测目标(即传入fit的y参数)就是输入本身,你当前代码将训练集特征作为输入x,却把测试集特征作为目标y传入,两者样本量差距极大,无法配对训练。
  • *数据基数(data cardinality)*在这里指代单个数组包含的样本总数量,也就是张量第一维的长度。Keras训练时要求所有输入数组、所有目标数组的第一维长度必须完全一致,否则无法按批次抽取配对样本。
  • 修正后的fit调用示例如下,注意测试集需要作为验证集传入validation_data参数:
hist = model.fit(
    x=[m_cat_train, m_num_train],
    y=[m_cat_train, m_num_train],  # 重构目标和输入一致
    batch_size=16,
    epochs=16,
    verbose=1,
    validation_data=([m_cat_test, m_num_test], [m_cat_test, m_num_test])
)
  • 额外注意:你给数值输出分支配置accuracy指标没有实际意义,accuracy是分类任务专用指标,数值分支是回归任务(用MSE损失),该指标计算结果不具备参考价值。
相关问题解答

1. 融合双输出指标的EarlyStopping实现

Keras原生EarlyStopping仅支持监听单个指标,要实现双指标融合判断停止条件,需要继承原生EarlyStopping类自定义回调,在每个epoch结束时按自定义规则计算两个指标的融合分数,再判断是否触发早停。示例代码如下(融合规则可根据任务调整,示例为分类准确率和数值MSE加权计算综合得分,分数越高代表模型效果越好):

from keras.callbacks import EarlyStopping
import numpy as np

class CombinedMetricES(EarlyStopping):
    def __init__(self, cat_weight=0.5, num_weight=0.5,
                 min_delta=0, patience=5, verbose=1,
                 baseline=None, restore_best_weights=True, start_from_epoch=0):
        super().__init__(
            min_delta=min_delta, patience=patience, verbose=verbose,
            mode="max", baseline=baseline, restore_best_weights=restore_best_weights,
            start_from_epoch=start_from_epoch
        )
        self.cat_w = cat_weight
        self.num_w = num_weight

    def get_monitor_value(self, logs):
        logs = logs or {}
        # 取两个分支的验证集指标
        cat_acc = logs.get("val_cat_output_accuracy", 0)
        num_mse = logs.get("val_num_output_mse", np.inf)
        # 融合规则:分类准确率权重*acc - 数值MSE权重*mse,统一为越大越优的分数
        combined_score = self.cat_w * cat_acc - self.num_w * num_mse
        return combined_score

# 调用示例
early_stop = CombinedMetricES(cat_weight=0.6, num_weight=0.4, patience=3)
# 训练时传入callbacks参数即可
# hist = model.fit(..., callbacks=[early_stop])

2. monitor设为val_accuracy时的监听对象

当多输出模型的EarlyStopping监听参数设置为不带输出名前缀的通用val_accuracy/accuracy时,Keras默认匹配输出列表中第一个输出的对应准确率指标,也就是你模型中最先定义的分类输出cat_output的验证准确率。

多输出模型的指标命名规则为{输出层名称}_{指标名},验证集指标会额外增加val_前缀,你当前模型实际存在的准确率指标为cat_output_accuracy、val_cat_output_accuracy,数值分支的准确率指标无实际参考价值。建议写全带输出层名称的完整指标名,避免匹配错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 18:45:39