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

训练LSTM模型时遇predictions must be <=1错误,求解决方案

问题:DKT模型训练时指标相关错误排查与解决

模型结构

inputs = tf.keras.Input(shape=(None, nb_features), name='inputs')

x = tf.keras.layers.Masking(mask_value=data.MASK_VALUE)(inputs)

x = tf.keras.layers.LSTM(hidden_units,
                         return_sequences=True,
                         dropout=dropout_rate)(x)

dense = tf.keras.layers.Dense(nb_skills, activation='sigmoid')
outputs = tf.keras.layers.TimeDistributed(dense, name='outputs')(x)

模型概况

Model: "DKTModel"
inputs (InputLayer)         [(None, None, 80)]        0
masking (Masking)           (None, None, 80)          0
lstm (LSTM)                 (None, None, 100)         72400
outputs (TimeDistributed)   (None, None, 40)          4040
=================================================================
Total params: 76,440
Trainable params: 76,440
Non-trainable params: 0

编译与fit函数实现

def compile(self, optimizer, metrics=None):
    
     def custom_loss(y_true, y_pred):
        y_true, y_pred = data.get_target(y_true, y_pred)
        return tf.keras.losses.binary_crossentropy(y_true, y_pred)
    
     super(DKTModel, self).compile(
        loss=custom_loss,
        optimizer=optimizer,
        metrics=metrics,
        experimental_run_tf_function=False)


def fit (self,
        dataset,
        epochs=1,
        verbose=1,
        callbacks=None,
        validation_data=None,
        shuffle=True,
        initial_epoch=0,
        steps_per_epoch=None,
        validation_steps=None,
        validation_freq=1):

  return super(DKTModel, self).fit(x=dataset, epochs=epochs,verbose=verbose, callbacks=callbacks, validation_data=validation_data, shuffle=shuffle, initial_epoch=initial_epoch, steps_per_epoch=steps_per_epoch, validation_steps=validation_steps, validation_freq=validation_freq)

错误日志

使用sigmoid激活时的错误

2 root error(s) found.(0) INVALID_ARGUMENT:  assertion failed: [predictions must be <= 1] [Condition x <= y did not hold element-wise:] [x (Sum_5:0) = ] [[[19.462822][19.5533848][19.5251656]]...] [y (Cast_11/x:0) = ] [1] [[{{node assert_less_equal/Assert/AssertGuard/Assert}}]][[assert_less_equal_2/Assert/AssertGuard/pivot_f/_122/_201]](1) INVALID_ARGUMENT:  assertion failed: [predictions must be <= 1] [Condition x <= y did not hold element-wise:] [x (Sum_5:0) = ] [[[19.462822][19.5533848][19.5251656]]...] [y (Cast_11/x:0) = ] [1][[{{node assert_less_equal/Assert/AssertGuard/Assert}}]]
0 successful operations.
0 derived errors ignored. [Op:__inference_train_function_7560]

改为softmax激活时的错误

2 root error(s) found.(0) INVALID_ARGUMENT:  assertion failed: [predictions must be <= 1] [Condition x <= y did not hold element-wise:] [x (Sum_5:0) = ] [[[0.99999994][1][1]]...] [y (Cast_11/x:0) = ] [1][{{node assert_less_equal/Assert/AssertGuard/Assert}}]]
 [[broadcast_weights_2/assert_broadcastable/AssertGuard/pivot_f/_58/_101]](1) INVALID_ARGUMENT:  assertion failed: [predictions must be <= 1] [Condition x <= y did not hold element-wise:] [x (Sum_5:0) = ] [[[0.99999994][1][1]]...] [y (Cast_11/x:0) = ] [1]
 [[{{node assert_less_equal/Assert/AssertGuard/Assert}}]]
0 successful operations.
0 derived errors ignored. [Op:__inference_train_function_7568]

当前指标设置

student_model.compile(
        optimizer=optimizer,
        metrics=[
            metrics.AUC(),
            metrics.Precision(),
            metrics.Recall()
        ])

解决建议

核心问题

自定义loss中通过data.get_target()对原始的y_true和y_pred做了维度/格式转换(比如从多技能的时序输出提取单技能的二分类目标),但内置指标直接使用原始输出计算,导致维度不匹配或数值不符合指标要求(比如Precision要求预测值是0-1的概率,但未处理的输出是多维度的,触发断言检查)。

具体解决步骤

  1. 为指标添加和loss一致的预处理逻辑
    包装内置指标,让它们先调用data.get_target()处理输入,再计算指标:

    def wrap_metric(metric_obj):
        def wrapped_metric(y_true, y_pred):
            # 和loss一样处理目标与预测值
            y_true_processed, y_pred_processed = data.get_target(y_true, y_pred)
            return metric_obj(y_true_processed, y_pred_processed)
        return wrapped_metric
    
    # 编译时使用包装后的指标
    student_model.compile(
        optimizer=optimizer,
        metrics=[
            wrap_metric(tf.keras.metrics.AUC()),
            wrap_metric(tf.keras.metrics.Precision()),
            wrap_metric(tf.keras.metrics.Recall())
        ]
    )
    
  2. 恢复正确的激活函数
    DKT是多标签分类任务(每个技能独立预测是否掌握),输出层必须用sigmoid激活,softmax适用于多分类(每个时间步只能选一个类别),不符合DKT的任务设定,改回sigmoid即可。

  3. 确认data.get_target()的输出格式
    确保该函数返回的y_true和y_pred是二分类格式:形状应为(None, None, 1)或(None, None),数值为0或1(y_true)、0-1之间的概率(y_pred),这样指标才能正确计算。

  4. 临时排查方案(可选)
    如果包装后仍有问题,可以尝试自定义指标类,跳过Keras内置的断言检查,或者先使用简单的指标(如tf.keras.metrics.BinaryAccuracy())验证模型是否能正常运行,再逐步添加复杂指标。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 16:18:33