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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 03:51:01