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

将LSTM全连接层上一时间步预测作为下一时间步额外输入的问题

解决LSTM将上一步Dense层预测结果作为下一步额外输入的问题

问题核心

你的代码未将上一步预测结果prev_pred与注意力输出context融合后传入LSTM,且初始化prev_pred时误用tf.keras.Input,这是拼接/Add操作报错的主要原因。以下是针对性修正方案:

变量维度对齐说明

根据你给出的维度:

  • context:需确保为(None, 1, 64)(若one_step_attention输出为2D(None,64),需扩展维度)
  • prev_pred:保持(None, 1, 30)
  • 融合后输入LSTM的张量需为(None, 1, 64+30)(拼接方案)或特征数统一后的同维度张量(相加方案)

修改后的完整代码

# 初始化prev_pred:直接用tf.zeros生成对应形状,避免Input层导致的维度冲突
prev_pred = tf.zeros((tf.shape(X)[0], 1, features))  # 形状(None, 1, 30)

for t in range(Ty):
    # Step 2.A: 获取当前时间步的context向量
    context = one_step_attention(X[:,t,:], cnn_model_input, s)
    
    # 确保context为3D张量(None,1,64),适配LSTM输入要求
    if len(context.shape) == 2:
        context = tf.expand_dims(context, axis=1)
    
    # 方案1:将context与prev_pred在特征维度拼接(推荐,无需调整特征数)
    combined_input = tf.concat([context, prev_pred], axis=-1)  # 形状(None,1,94)
    
    # 方案2:若要使用Add(),先将prev_pred投影到context的特征维度64
    # prev_pred_proj = Dense(64, activation='linear')(prev_pred)
    # combined_input = tf.add(context, prev_pred_proj)
    
    # 将融合后的输入传入LSTM,更新状态
    s, _, c = LSTM(n_s, return_state=True)(combined_input, initial_state=[s, c])
    
    # 生成当前时间步的预测结果
    out = Dense(features, activation='linear')(s)
    # 将out扩展为3D,作为下一步的prev_pred
    prev_pred = tf.expand_dims(out, axis=1)
    
    # 收集输出
    outputs.append(out)

# 创建模型
model = Model(inputs=[X_CNN_input, X_lstm_input, s0, c0], outputs=outputs)
return model

关键修正点

  1. 初始化修正:去掉tf.keras.Input初始化逻辑,直接用tf.zeros结合输入X的batch维度生成初始预测张量,避免Input层带来的图结构冲突。
  2. 维度统一:确保context和prev_pred的时间步维度均为1,保证拼接/相加操作的维度合法性。
  3. 输入融合:将融合后的张量传入LSTM,让模型同时接收注意力上下文和上一步预测信息。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 15:02:40