如何让Keras LSTM层接收双输入并实现时间序列对遗忘门的调控?
兄弟,这个报错我太熟了,本质是你给LSTM层传参的方式完全不对,而且自定义Cell的调用逻辑也没跟上Keras的规则,咱们一步步拆解解决:
第一步:搞懂报错的根源
你写的lstm_out = LSTM(32)(x, time_vector = m)犯了一个关键错误:Keras内置的LSTM层根本没有time_vector这个参数! 当你传了一个它不认识的参数时,Keras会默认把它当成initial_state(模型初始状态)来处理,但你的m是形状为(None,50,1)的时间序列,而LSTMCell的状态是两个形状为(None,32)的张量(隐藏状态和细胞状态),维度完全不匹配,所以直接炸出了这个报错。
第二步:正确实现带时间影响遗忘门的自定义LSTMCell
你需要把“时间序列影响遗忘门”的逻辑完整封装到自定义Cell里,并且确保它能正确接收两个输入(事件embedding和时间序列)。这里要注意:时间序列的每个步长是标量,需要先转换成和遗忘门同维度的权重,避免维度不匹配。
修改后的自定义Cell示例:
import tensorflow as tf from tensorflow.keras.layers import LSTMCell from tensorflow.keras import backend as K class TimeAffectedLSTMCell(LSTMCell): def __init__(self, units, **kwargs): super().__init__(units, **kwargs) # 加一个Dense层把时间标量映射成和遗忘门同维度的权重 self.time_weight_proj = tf.keras.layers.Dense(units, activation='sigmoid') def call(self, inputs, states, training=None): # inputs是列表:[当前步的事件embedding, 当前步的时间值] event_input = inputs[0] time_input = inputs[1] # 处理时间输入:计算遗忘权重,用+1避免除以0 time_weight = 1.0 / (K.cast(time_input, K.floatx()) + 1.0) # 扩展维度后映射到遗忘门的维度(units维) time_weight = K.expand_dims(time_weight, axis=-1) time_weight = self.time_weight_proj(time_weight) # 原有LSTM的计算逻辑 h_tm1 = states[0] # 上一步隐藏状态 c_tm1 = states[1] # 上一步细胞状态 z = K.dot(event_input, self.kernel) + K.dot(h_tm1, self.recurrent_kernel) z += self.bias i = self.recurrent_activation(z[:, :self.units]) f = self.recurrent_activation(z[:, self.units: self.units*2]) # 核心操作:把时间权重乘到遗忘门上 f = f * time_weight c = self.recurrent_activation(z[:, self.units*2: self.units*3]) o = self.recurrent_activation(z[:, self.units*3:]) c_t = f * c_tm1 + i * c h_t = o * self.activation(c_t) return h_t, [h_t, c_t]
第三步:用RNN层包裹自定义Cell,正确传递多输入
Keras内置LSTM层不支持自定义多输入,所以你需要用tf.keras.layers.RNN来包裹自己的Cell,并且把两个输入(事件embedding和时间序列)合并成列表传入:
修改后的模型构建代码:
from tensorflow.keras.layers import Input, Embedding, Masking, Dense, RNN from tensorflow.keras.models import Model import numpy as np max_seq_length = 50 embedding_length = 64 num_unique_event_symbols = 101 # 事件是1-100的整数,所以输入维度设为101 # Input 1: 事件类型序列 main_input = Input(shape=(max_seq_length,), dtype='int32', name='main_input') x = Embedding(output_dim=embedding_length, input_dim=num_unique_event_symbols, input_length=max_seq_length, mask_zero=True)(main_input) # Input 2: 时间间隔序列 auxiliary_input = Input(shape=(max_seq_length,1), dtype='float32', name='aux_input') m = Masking(mask_value=99999999.0)(auxiliary_input) # 用RNN层包裹自定义Cell,传入两个输入的列表 custom_cell = TimeAffectedLSTMCell(32) lstm_out = RNN(custom_cell)([x, m]) # 后续输出层保持不变 auxiliary_output = Dense(1, activation='sigmoid', name='aux_output')(lstm_out) x = Dense(64, activation='relu')(lstm_out) main_output = Dense(1, activation='sigmoid', name='main_output')(x) # 编译与训练 model = Model(inputs=[main_input, auxiliary_input], outputs=[main_output, auxiliary_output]) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'], loss_weights=[1., 0.2]) print(model.summary()) np.random.seed(21) # 假设你已准备好train_X1, train_X2, train_Y数据 model.fit([train_X1, train_X2], [train_Y, train_Y], epochs=1, batch_size=200)
额外注意事项
- Masking处理:确保时间序列的填充值被Masking层正确过滤,避免影响遗忘门的计算逻辑。
- 时间权重公式:你可以根据需求调整时间权重的计算方式,比如
(1/(time_input +1))或者其他单调递减函数,只要避免除以0即可。 - 维度匹配:自定义Cell会自动按时间步处理
(事件embedding, 时间值)的输入对,无需手动拆分序列。
这样修改后,你的模型就能正确把时间序列信息注入到遗忘门中,同时解决维度不匹配的报错问题。
内容的提问来源于stack exchange,提问作者Slyron
相关产品推荐
相关产品推荐

