Keras中如何通过共享权重的LSTM层处理多组输入序列?
实现方案
Keras中相同的层实例被多次调用时会自动共享权重,你需要的功能只要对应调整输入结构即可,以下是两种可直接运行的实现方案:
方案1:TimeDistributed包装实现(推荐,代码简洁,支持动态序列数量)
适合序列数量不固定、或者序列数量较大的场景,通过TimeDistributed自动对每组序列应用同一个LSTM层:
# 输入shape调整为(序列组数, 单序列最大长度, 1) sq_inpt = Input(shape=(self.num_sequences, self.MAX_SEQ_LEN, 1), name='sq_inpt') embed = Embedding(output_dim=100, input_dim=self.sequence_range + 1, mask_zero=True, name='sequence_embedding') # 经过Embedding后shape为 (batch_size, 序列组数, MAX_SEQ_LEN, 100) sq_embed = embed(sq_inpt) # 仅实例化一次LSTM层,后续全程复用实现权重共享 lstm_layer = LSTM(200, go_backwards=False) # TimeDistributed会自动对第二维度的每组序列独立应用LSTM层 lstm_embed = TimeDistributed(lstm_layer)(sq_embed) # 展平所有序列的LSTM输出,用于后续拼接 lstm_embed_flat = Flatten()(lstm_embed) # 后续逻辑和原模型保持一致 time_inpt = Input(shape=(1,), name='time_inpt') score_inpt = Input(shape=(1,), name='score_inpt') state_embed = Concatenate()([lstm_embed_flat, time_inpt, score_inpt]) state_embed = Dense(300, activation='elu', name='state_embed_1')(state_embed) state_embed = Dense(300, activation='elu', name='state_embed_2')(state_embed) output = Dense(1, name='output')(state_embed) model = Model(inputs=[sq_inpt, time_inpt, score_inpt], outputs=output)
方案2:固定序列数量手动循环实现
适合序列数量少且固定的场景,逻辑和你写的示意代码完全匹配:
# 假设固定有3组序列,输入shape为(3, MAX_SEQ_LEN, 1) sq_inpt = Input(shape=(3, self.MAX_SEQ_LEN, 1), name='sq_inpt') embed = Embedding(output_dim=100, input_dim=self.sequence_range + 1, mask_zero=True, name='sequence_embedding') sq_embed = embed(sq_inpt) lstm_layer = LSTM(200, go_backwards=False) lstm_embed = [] # 遍历每组序列,复用同一个LSTM层处理 for i in range(3): # 取出第i组序列的embedding结果 single_seq_embed = sq_embed[:, i, :, :] lstm_embed.append(lstm_layer(single_seq_embed)) # 拼接所有序列的LSTM输出 lstm_embed_flat = Concatenate()(lstm_embed) # 后续逻辑和原模型保持一致 time_inpt = Input(shape=(1,), name='time_inpt') score_inpt = Input(shape=(1,), name='score_inpt') state_embed = Concatenate()([lstm_embed_flat, time_inpt, score_inpt]) state_embed = Dense(300, activation='elu', name='state_embed_1')(state_embed) state_embed = Dense(300, activation='elu', name='state_embed_2')(state_embed) output = Dense(1, name='output')(state_embed) model = Model(inputs=[sq_inpt, time_inpt, score_inpt], outputs=output)
注意事项
- 两种方案都仅实例化了一次LSTM层,所有序列处理复用同一个层实例,天然实现权重共享
- 如果你的单序列长度不统一,把输入里的
MAX_SEQ_LEN改为None即可,LSTM本身支持动态长度输入
内容的提问来源于stack exchange,提问作者Ashwin Kumar
相关产品推荐
相关产品推荐

