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

如何基于LSTM/GRU结合Keras Attention层实现二分类编码器-解码器架构?

解决Keras内置Attention层用于二分类的编码器-解码器架构问题

我明白你的困惑——官方的Attention示例大多是针对序列到序列翻译的,而你需要用它来做二分类,确实得调整架构。下面我会一步步帮你把代码补全,解释每个部分的作用:

首先,先明确核心逻辑:对于二分类任务,我们不需要解码器逐步生成序列,只需要利用注意力机制从编码器的所有时间步隐藏状态中提取关键信息,再结合解码器的状态得到最终的分类特征。

完整实现代码

import tensorflow as tf

# 输入层:形状(batch_size, 时间步长100, 特征维度50)
x = tf.keras.Input((100, 50))

# ---------------------- 编码器部分 ----------------------
# 使用GRU,同时返回所有时间步的隐藏状态(供Attention使用)和最后一个时间步的状态(作为解码器初始状态)
encoder_gru = tf.keras.layers.GRU(32, return_sequences=True, return_state=True)
encoder_hidden_states, encoder_final_state = encoder_gru(x)
# encoder_hidden_states: (batch_size, 100, 32) 所有时间步的隐藏状态
# encoder_final_state: (batch_size, 32) 编码器最后一个状态,传给解码器做初始状态

# ---------------------- 解码器 + Attention部分 ----------------------
# 解码器的初始输入:因为我们不需要生成序列,用一个全零的单步向量即可(形状(batch_size, 1, 32))
# 也可以换成可训练的初始向量,效果可能更好,后面会说明
decoder_initial_input = tf.keras.layers.Lambda(lambda x: tf.zeros_like(x[:, 0:1, :]))(encoder_hidden_states)

# 解码器GRU:只运行一步,返回该步的输出和状态
decoder_gru = tf.keras.layers.GRU(32, return_sequences=True, return_state=True)
decoder_outputs, _ = decoder_gru(decoder_initial_input, initial_state=encoder_final_state)
# decoder_outputs: (batch_size, 1, 32) 解码器单步输出,作为Attention的query

# 内置Attention层:输入是[query, value],这里query是解码器输出,value是编码器所有隐藏状态
attention_layer = tf.keras.layers.Attention()
context_vector = attention_layer([decoder_outputs, encoder_hidden_states])
# context_vector: (batch_size, 1, 32) 注意力加权后的上下文向量

# 拼接上下文向量和解码器输出,融合两种信息
concat_features = tf.keras.layers.Concatenate(axis=-1)([context_vector, decoder_outputs])
# 展平成2D张量,供分类层使用
flattened_features = tf.keras.layers.Flatten()(concat_features)

# ---------------------- 分类部分 ----------------------
z = tf.keras.layers.Dense(1, activation='sigmoid')(flattened_features)

# 构建模型
model = tf.keras.Model(inputs=x, outputs=z)
model.summary()

关键部分解释

  1. 编码器的状态返回
    我们给GRU加上return_state=True,这样能拿到编码器最后一个时间步的隐藏状态,把它作为解码器的初始状态,让解码器能继承编码器的全局信息。

  2. 解码器的初始输入
    因为不需要生成序列,解码器只需要运行一步,所以初始输入用全零向量就足够。如果你想让初始输入更灵活,可以换成可训练的向量:

    # 创建一个可训练的初始输入向量
    decoder_initial_input = tf.keras.layers.Lambda(
        lambda _: tf.tile(tf.Variable(tf.random.normal((1,1,32))), [tf.shape(_)[0], 1, 1])
    )(x)
    
  3. Attention层的输入要求
    Keras的Attention层要求输入都是3D张量((batch_size, timesteps, features)),所以解码器的输出是(batch_size,1,32)(单步),编码器的隐藏状态是(batch_size,100,32)(所有时间步),这样才能计算注意力权重。

  4. 上下文向量的使用
    我们把注意力得到的上下文向量和解码器输出拼接,是为了同时利用注意力提取的关键时序信息,和解码器继承的编码器全局信息,让分类特征更全面。

简化版(如果不需要完整解码器)

如果你觉得解码器的GRU有点多余,也可以直接用编码器的最后状态作为query,计算注意力后直接分类:

import tensorflow as tf

x = tf.keras.Input((100, 50))
encoder_gru = tf.keras.layers.GRU(32, return_sequences=True, return_state=True)
encoder_hidden_states, encoder_final_state = encoder_gru(x)

# 把编码器最后状态变成3D张量(符合Attention输入要求)
query = tf.keras.layers.Lambda(lambda x: tf.expand_dims(x, axis=1))(encoder_final_state)
# 计算注意力上下文
context_vector = tf.keras.layers.Attention()([query, encoder_hidden_states])
# 展平后分类
flattened = tf.keras.layers.Flatten()(context_vector)
z = tf.keras.layers.Dense(1, activation='sigmoid')(flattened)

model = tf.keras.Model(inputs=x, outputs=z)

这个简化版更轻量,适合对模型复杂度有要求的场景,效果也不会差太多。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 09:57:29