欢迎光临
我们一直在努力

让模型学会“回头看”:Seq2Seq 与 Attention 机制的通俗图解与代码实战【NLP系列第三篇】

让模型学会“回头看”:Seq2Seq 与 Attention 机制的通俗图解与代码实战

1. 引子:从"单兵作战"到"双人配合"

前两篇我们聊了词向量和 RNN/LSTM/GRU。但仔细一想,之前讲的场景都是输入一个序列,输出一个结果——比如情感分类,读完整句话判断个正负面。

但 NLP 里很多任务是输入是序列,输出也是序列,而且长度还不一样:

  • 机器翻译:输入"我爱你"(3个字),输出"I love you"(3个词)
  • 文本摘要:输入一篇 5000 字的文章,输出 200 字的摘要
  • 语音识别:输入一段音频,输出一串文字

一个 RNN 搞不定这种事。你需要两个 RNN 配合:

  • 第一个负责读完源句,把意思"压缩"成一个向量
  • 第二个负责拿着这个向量,逐字生成目标句子

这就是 Seq2Seq(Encoder-Decoder) 架构。


2. Seq2Seq:编码器-解码器架构

Seq2Seq 的命名很直白:输入一个序列(Sequence),输出另一个序列(Sequence)。它由两个核心部分组成。

### 2.1 Encoder:把源句"压缩"成向量

Encoder 就是一个 RNN(LSTM/GRU 也可以),它会依次读入源句中的每个词,最后一个时间步的隐状态就被当作上下文向量(Context Vector)——可以理解为整个句子的语义浓缩。

这里有两个工程细节值得注意:

  • Encoder 的隐状态怎么取? 可以只取最后一个时间步的

    h

    n

    h_n

    hn,也可以对所有时间步的隐状态做平均池化或最大池化。最常见的做法是取最后一个时间步的输出。

  • 双向还是单向? 双向 RNN 能同时看到前后文,比单向捕捉的语义更完整。Encoder 通常用双向 LSTM/GRU。
  • 2.2 Decoder:自回归生成

    Decoder 是另一个 RNN,它的任务是从上下文向量中"解压"出目标序列。

    生成过程是自回归(Autoregressive) 的——当前时间步的输出会作为下一个时间步的输入:

  • 初始状态:Decoder 的初始隐状态来自 Encoder 的上下文向量
  • 起始符:第一个输入是一个特殊标记 <sos>(Start of Sentence),告诉模型"该开始生成
  • 逐字生成:每一步根据当前隐状态和上一步的输出,预测下一个 token
  • 结束符:当模型生成 <eos>(End of Sentence)时停下来
  • # 自回归生成伪代码
    def generate(encoder_hidden, sos_token, max_len=50):
    # Decoder 初始状态 = Encoder 最后时刻隐状态
    decoder_hidden = encoder_hidden
    # 第一个输入是 <sos>
    decoder_input = sos_token
    outputs = []

    for t in range(max_len):
    # 当前步预测
    output, decoder_hidden = decoder_step(decoder_input, decoder_hidden)
    outputs.append(output)

    # 用当前步的输出作为下一步的输入
    decoder_input = output.argmax(dim=1)

    # 如果预测到 <eos> 就提前停止
    if decoder_input == eos_token:
    break

    return outputs

    注意:推理时用上一步的预测结果作为下一步输入,这叫"贪婪解码"。训练时一般用"Teacher Forcing"——直接把真实上一步的 token 喂给下一步,加速收敛。

    2.3 致命缺陷:信息瓶颈

    Seq2Seq 的结构很简洁,但有个硬伤——

    无论输入多长,Encoder 都要把所有信息压缩进一个固定长度的上下文向量。

    短句还好说,但句子一长(比如翻译一段 50 个词的段落),这个向量根本装不下全部信息。前面的内容被后面的冲刷掉,Decoder 拿到的是一个"缩水版"的语义。这就是 Seq2Seq 的信息瓶颈。

    而且 Decoder 在每个时间步看到的都是同一个向量——生成主语时它在看这个向量,生成谓语时它还在看这个向量。但主语和谓语依赖的源句信息显然是不同的。

    怎么解决?给 Decoder 装个"探照灯"。


    3. 注意力机制:让 Decoder 学会"回头看"

    3.1 核心思想

    传统 Seq2Seq 中,Decoder 只依赖 Encoder 最后一个隐状态。

    注意力机制(Attention Mechanism)的改进很直观:不再只依赖那个压缩后的固定向量,而是让 Decoder 在每一步生成时,都能"回头看"一眼 Encoder 对所有输入词的隐状态,然后动态选择当前最该关注哪些位置。

    换成人类翻译的例子来理解:

    把 “I love you” 翻译成中文

    • 生成"我"的时候,最该关注的是源句中的 “I”
    • 生成"爱"的时候,最该关注的是 “love”
    • 生成"你"的时候,最该关注的是 “you”

    注意力机制就是让模型学会这种"对齐"关系。

    3.2 四步流程详解

    注意: 注意力机制主要分为两大流派,本笔记及配图基于 Luong Attention (2015) 的实现逻辑编写。

    • Bahdanau Attention (Additive): “先看后动”。使用上一时刻的状态

      s

      t

      1

      s_{t−1}

      st1​ 计算注意力,得到上下文

      c

      t

      c_t

      ct​ 后再更新当前状态。

    • Luong Attention (Multiplicative): “先动后看”(本笔记采用)。RNN 先根据输入更新出当前时刻的原始状态

      s

      t

      s_t

      st​ ,再用这个

      s

      t

      s_t

      st​ 去计算注意力并融合。

    两者的核心区别在于计算注意力分数时使用的是

    s

    t

    1

    s_{t−1}

    st1​ 还是

    s

    t

    s_t

    st​ 。 请阅读时留意这一差异,以免与部分经典教材混淆。

    注意力计算一共四步,每一步都有明确的数学操作:

    符号说明:

    • Encoder 输出:

      h

      1

      ,

      h

      2

      ,

      ,

      h

      T

      x

      h_1, h_2, \\dots, h_{T_x}

      h1,h2,,hTx(每个

      h

      i

      h_i

      hi 对应源句的一个词)

    • Decoder 当前隐状态:

      s

      t

      s_{t}

      st(rnn 基于上一步解码状态 + 当前时间步的输入)


    第一步:计算注意力分数

    对于当前要输出的第

    t

    t

    t 个词,计算

    s

    t

    s_{t}

    st 与每个

    h

    i

    h_i

    hi 的相关性分数:

    e

    t

    ,

    i

    =

    align

    (

    s

    t

    ,

    h

    i

    )

    e_{t,i} = \\text{align}(s_{t}, h_i)

    et,i=align(st,hi)

    align 是评分函数,常见有 Dot、General、Concat 三种(下节讲)。


    第二步:softmax 归一化

    把分数转成概率分布,表示每个输入位置对当前输出的重要程度:

    α

    t

    ,

    i

    =

    exp

    (

    e

    t

    ,

    i

    )

    k

    =

    1

    T

    x

    exp

    (

    e

    t

    ,

    k

    )

    \\alpha_{t,i} = \\frac{\\exp(e_{t,i})}{\\sum_{k=1}^{T_x} \\exp(e_{t,k})}

    αt,i=k=1Txexp(et,k)exp(et,i)

    所有

    α

    t

    ,

    i

    \\alpha_{t,i}

    αt,i 加起来等于 1。


    第三步:加权求和得到上下文向量

    用注意力权重对 Encoder 隐状态做加权平均,得到当前步最该关注的上下文:

    c

    t

    =

    i

    =

    1

    T

    x

    α

    t

    ,

    i

    h

    i

    c_t = \\sum_{i=1}^{T_x} \\alpha_{t,i} \\cdot h_i

    ct=i=1Txαt,ihi

    这个

    c

    t

    c_t

    ct 就是为当前输出位置"量身定制"的上下文——和传统 Seq2Seq 的固定向量不同,每步的

    c

    t

    c_t

    ct 都不同。


    第四步:结合上下文预测输出

    把上下文向量

    c

    t

    c_t

    ct 和 Decoder 当前隐状态

    s

    t

    s_t

    st 拼接起来,通过一个线性层 + softmax 预测下一个词:

    s

    ~

    t

    =

    tanh

    (

    W

    [

    s

    t

    ;

    c

    t

    ]

    )

    \\tilde{s}_t = \\tanh(W[s_t; c_t])

    s~t=tanh(W[st;ct])

    P

    (

    y

    t

    y

    <

    t

    ,

    x

    )

    =

    softmax

    (

    V

    s

    ~

    t

    )

    P(y_t \\mid y_{<t}, x) = \\text{softmax}(V \\tilde{s}_t)

    P(yty<t,x)=softmax(Vs~t)

    这一步做完,Decoder 就拿到了当前步最需要的源句信息,预测的准确率自然比只看固定向量高得多。

    Decoder 每步的完整流程总结:

    更新状态

    s

    t

    s_t

    st → 查源句重点

    c

    t

    c_t

    ct → 融合状态与上下文 → 预测下一个词


    4. 评分函数:怎么算"相关性"

    第一步中计算分数

    e

    t

    ,

    i

    e_{t,i}

    et,i 的 align 函数,常见三种实现:

    评分函数公式特点
    Dot(点积)

    e

    =

    s

    t

    h

    i

    e = s_{t} \\cdot h_i

    e=sthi

    最简单,无额外参数,要求

    s

    s

    s

    h

    h

    h 维度相同

    General(通用点积)

    e

    =

    s

    t

    T

    W

    a

    h

    i

    e = s_{t}^T W_a h_i

    e=stTWahi

    中间插一个可学习的

    W

    a

    W_a

    Wa,维度不同也能适配

    Concat(拼接/加性)

    e

    =

    v

    a

    T

    tanh

    (

    W

    a

    [

    s

    t

    1

    ;

    h

    i

    ]

    )

    e = v_a^T \\tanh(W_a [s_{t – 1}; h_i])

    e=vaTtanh(Wa[st1;hi])

    即 Bahdanau 注意力,参数最多,表达能力最强
    • Dot 最快但最"笨",维度一高点积值容易偏大,softmax 梯度消失
    • General 是灵活度和速度的折中
    • Concat(也叫加性注意力)最灵活,但计算开销也最大

    在 Transformer 出现之前,Concat(Bahdanau)是最主流的实现。


    5. 从交叉注意力到自注意力

    5.1 交叉注意力

    上面讲的注意力,Q(Query)来自 Decoder,K(Key)和 V(Value)都来自 Encoder。这种 Q 和 KV 来自不同地方的注意力,叫做 交叉注意力(Cross-Attention),也叫 Encoder-Decoder 注意力。

    它的作用是:让 Decoder 在生成时对齐到源句的对应位置。

    5.2 自注意力

    那如果 Q、K、V 都来自同一个序列呢?这就是 自注意力(Self-Attention)。

    举个例子,看这个句子:

    The animal didn’t cross the street because it was too tired.

    “it” 指的是谁?人类一看就知道是 “The animal”。但计算机怎么知道?自注意力就是让序列中的每个 token 去"关注"序列中的其他所有 token,然后把相关的信息聚合到自己身上。

    当模型处理 “it” 这个位置时,它会通过自注意力给 “The animal” 分配很高的权重,给 “street” 分配很低的权重。这样 “it” 的最终表示中就包含了 “The animal” 的语义信息。

    QKV 三个角色的类比:

    • Query:当前位置的"需求"——“我想要找什么信息”
    • Key:各个位置的"索引"——“我有什么信息可以提供”
    • Value:各个位置的"内容"——“我能提供的信息具体是什么”

    Q 和 K 算相似度决定关注谁,然后从 V 中把对应的内容提取出来。

    5.3 缩放点积注意力

    自注意力用的是 缩放点积(Scaled Dot-Product) 公式:

    Attention

    (

    Q

    ,

    K

    ,

    V

    )

    =

    softmax

    (

    Q

    K

    T

    d

    k

    )

    V

    \\text{Attention}(Q, K, V) = \\text{softmax}\\left(\\frac{QK^T}{\\sqrt{d_k}}\\right)V

    Attention(Q,K,V)=softmax(dk

    QKT)V

    这个公式相比 Bahdanau 加性注意力有两个优势:

  • 矩阵乘法一次性完成:Q 和 K 的维度是 (seq_len, d_k),一次

    Q

    K

    T

    QK^T

    QKT 就算出了所有位置两两之间的分数,不用逐个拼接。这就是全并行计算的来源。

  • 缩放因子

    d

    k

    \\sqrt{d_k}

    dk

    :维度越高,点积值越大,softmax 后的梯度越容易消失。除

    d

    k

    \\sqrt{d_k}

    dk

    让分数落在合理的数值范围。

  • 5.4 掩码自注意力

    Decoder 的自注意力还有一个特殊限制:当前位置不能看到未来的词。

    比如在机器翻译中,生成 “love” 这个词时,Decoder 可以看到 <sos> 和 “I”,但不能看到还没生成的 “you”。否则就是作弊。

    实现方式很简单:在计算分数后,用一个下三角矩阵把未来位置设为

    -\\infty

    ,这样 softmax 后它们的权重就是 0。

    # 下三角掩码矩阵(seq_len=5)
    [ 1 0 0 0 0
    1 1 0 0 0
    1 1 1 0 0
    1 1 1 1 0
    1 1 1 1 1 ]

    这种"只能看过去和当前,不能看未来"的自注意力,叫掩码自注意力(Masked Self-Attention)。

    5.5 交叉注意力 vs 自注意力

    对比维度交叉注意力自注意力
    Q 来源 Decoder 隐状态 输入序列本身
    KV 来源 Encoder 所有隐状态 输入序列本身
    作用 输入输出对齐 建立序列内部依赖
    计算范围 两个序列之间 单个序列内部

    6. PyTorch 代码实战

    6.1 缩放点积自注意力

    下面是 Transformer 最核心的模块——缩放点积自注意力的完整 PyTorch 实现,每行代码都标注了张量维度的变化:

    import torch
    import torch.nn as nn
    import torch.nn.functional as F

    class ScaledDotProductSelfAttention(nn.Module):
    def __init__(self, embed_dim):
    """
    :param embed_dim: 输入输出嵌入维度
    """

    super().__init__()
    self.embed_dim = embed_dim
    # Q、K、V 三个可学习的投影矩阵
    self.W_q = nn.Linear(embed_dim, embed_dim) # 输入 embed_dim → 输出 embed_dim
    self.W_k = nn.Linear(embed_dim, embed_dim)
    self.W_v = nn.Linear(embed_dim, embed_dim)

    def forward(self, x, mask=None):
    """
    :param x: 输入序列 (batch_size, seq_len, embed_dim)
    :param mask: 可选掩码, True 表示需要 mask 掉该位置
    :return: (output, attention_weights)
    """

    batch_size, seq_len, embed_dim = x.shape

    # 1. 对每个 token 投影得到 Q、K、V
    # 输入: (batch, seq_len, embed_dim)
    # 输出: (batch, seq_len, embed_dim)
    q = self.W_q(x)
    k = self.W_k(x)
    v = self.W_v(x)

    # 2. 计算注意力分数: (Q @ K^T) / sqrt(d_k)
    # q: (batch, seq_len, embed_dim)
    # k.transpose: (batch, embed_dim, seq_len)
    # scores: (batch, seq_len, seq_len) ← 每个位置对所有位置的分数
    scores = torch.bmm(q, k.transpose(1, 2)) / torch.sqrt(
    torch.tensor(embed_dim, dtype=torch.float32)
    )

    # 3. 如果有掩码,把未来位置设为 -inf(softmax 后就是 0)
    if mask is not None:
    # mask shape: (batch, seq_len), True = 需要 mask
    scores = scores.masked_fill(mask.unsqueeze(1), float('inf'))

    # 4. softmax 归一化得到注意力权重
    # (batch, seq_len, seq_len) — 每行表示当前 token 对其他 token 的关注度
    attn_weights = F.softmax(scores, dim=1)

    # 5. 加权求和得到输出
    # (batch, seq_len, seq_len) @ (batch, seq_len, embed_dim)
    # = (batch, seq_len, embed_dim) ← 输出和输入形状一致
    output = torch.bmm(attn_weights, v)

    return output, attn_weights

    使用示例:

    # batch_size=2, seq_len=5, embed_dim=128
    x = torch.randn(2, 5, 128)

    self_attn = ScaledDotProductSelfAttention(embed_dim=128)
    output, attn = self_attn(x)

    print(f"输入: {x.shape}") # torch.Size([2, 5, 128])
    print(f"输出: {output.shape}") # torch.Size([2, 5, 128]) ← 输入输出维度不变
    print(f"权重: {attn.shape}") # torch.Size([2, 5, 5]) ← 5×5 注意力矩阵

    掩码自注意力(Decoder 用):

    只需要在上面代码的 forward 中加两行生成下三角掩码:

    class MaskedSelfAttention(ScaledDotProductSelfAttention):
    def forward(self, x):
    batch_size, seq_len, _ = x.shape

    q = self.W_q(x)
    k = self.W_k(x)
    v = self.W_v(x)

    scores = torch.bmm(q, k.transpose(1, 2)) / torch.sqrt(
    torch.tensor(self.embed_dim, dtype=torch.float32)
    )

    # === 下三角掩码:只保留当前位置及之前的位置 ===
    # tril 生成下三角为 1 的矩阵,未来位置设为 -inf
    mask = torch.tril(torch.ones(seq_len, seq_len, device=x.device))
    # mask = 0 的位置就是未来位置,设为 -inf
    scores = scores.masked_fill(mask == 0, float('inf'))

    attn_weights = F.softmax(scores, dim=1)
    output = torch.bmm(attn_weights, v)

    return output, attn_weights

    6.2 两种注意力对比

    对比维度Bahdanau 加法注意力缩放点积注意力
    适用场景 传统 RNN Encoder-Decoder Transformer
    计算方式 每个 Key 拼接 Query → 线性层算分数 Q 点积 K^T → 直接算
    复杂度 逐位置计算,串行 矩阵乘法一次性完成,全并行
    效率 较慢
    使用地方 传统 NMT 解码器 Transformer 编码器 / 解码器

    6.3 多头自注意力

    上面实现的单头自注意力有一个局限:所有 attention 都共享同一组 QKV 投影。但一句话里往往同时包含多种语义关系——句法、词义、指代等等。单头注意力的"视野"是有限的。

    多头注意力(Multi-Head Attention) 的做法很直接:用多组独立的 QKV 投影(多个"头"),每个头学习不同类型的注意力依赖,最后把各头的输出拼回去:

    • 某些头关注句法依赖(主谓宾)
    • 某些头关注共指关系(“it” → “animal”)
    • 某些头关注长距离语义关联

    import torch
    import torch.nn as nn
    import torch.nn.functional as F

    class MultiHeadSelfAttention(nn.Module):
    def __init__(self, dim, num_heads):
    """
    :param dim: 输入输出维度
    :param num_heads: 注意力头数
    """

    super().__init__()
    assert dim % num_heads == 0, "dim 必须能被 num_heads 整除"

    self.dim = dim
    self.num_heads = num_heads
    self.head_dim = dim // num_heads # 每个头的维度 = 总维度 / 头数

    # 1. 定义 Q, K, V 投影矩阵(注意:不是每个头单独定义,而是一次性投影到完整 dim)
    # 稍后在 forward 中拆分为多个头
    self.W_q = nn.Linear(dim, dim, bias=False)
    self.W_k = nn.Linear(dim, dim, bias=False)
    self.W_v = nn.Linear(dim, dim, bias=False)

    # 2. 多头拼接后的输出投影
    self.out_proj = nn.Linear(dim, dim, bias=False)

    def forward(self, x):
    """
    :param x: (batch_size, seq_len, dim)
    :return: (output, attn_weights)
    """

    batch_size, seq_len, dim = x.shape

    # 1. 线性投影得到 Q, K, V
    # 输入: (batch, seq_len, dim)
    # 输出: (batch, seq_len, dim)
    q = self.W_q(x)
    k = self.W_k(x)
    v = self.W_v(x)

    # 2. 拆分多头:把最后 1 维拆成 (num_heads, head_dim)
    # view 后: (batch, seq_len, num_heads, head_dim)
    # 转置后: (batch, num_heads, seq_len, head_dim)
    # 这样每个头独立计算注意力,互不干扰
    q = q.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
    k = k.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
    v = v.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

    # 3. 所有头并行计算缩放点积注意力
    # (batch, num_heads, seq_len, seq_len)
    scores = torch.matmul(q, k.transpose(2, 1)) / torch.sqrt(
    torch.tensor(self.head_dim, dtype=torch.float32)
    )

    # 4. Softmax
    attn_weights = F.softmax(scores, dim=1) # (batch, num_heads, seq_len, seq_len)

    # 5. 加权求和
    # (batch, num_heads, seq_len, head_dim)
    context = torch.matmul(attn_weights, v)

    # 6. 拼接多头:转置回 (batch, seq_len, num_heads, head_dim),再合并后两维
    # → (batch, seq_len, dim)
    context = context.transpose(1, 2).contiguous().view(batch_size, seq_len, dim)

    # 7. 最终线性投影
    output = self.out_proj(context) # (batch, seq_len, dim)

    return output, attn_weights

    维度变化之旅:

    步骤维度变化说明
    输入 x (batch, seq_len, dim) 原始输入
    QKV 投影 同上 保持维度不变
    view + 转置 (batch, num_heads, seq_len, head_dim) 拆分为多个头
    scores (batch, num_heads, seq_len, seq_len) 每个头计算各自的注意力矩阵
    context (batch, num_heads, seq_len, head_dim) 每个头加权求和的结果
    转置 + view (batch, seq_len, dim) 所有头拼接回原维度
    输出 (batch, seq_len, dim) 最终投影输出

    关键点:多头注意力的计算量和单头几乎一样(总维度不变),但表达能力更强——因为每个头可以关注不同的子空间。


    7. 总结与下篇预告

    本文核心脉络:

  • Seq2Seq:Encoder 压缩源句为固定向量,Decoder 自回归生成——简单但存在信息瓶颈
  • 注意力机制:Decoder 每一步动态关注源句不同位置,四步流程(算分→归一化→加权→融合)
  • 自注意力:QKV 来自同一序列,矩阵乘法一次算出所有位置两两关系,实现全并行
  • 掩码自注意力:Decoder 专用,防止看到未来信息
  • 多头注意力:多组 QKV 并行,每个头关注不同的子空间(句法、共指、语义),拼接后输出
  • 理解这些,Transformer 的骨架就已经搭好了。剩下的就是在此基础上加:

    • 位置编码(自注意力没有位置感,需要额外注入)
    • 前馈神经网络 + 残差连接 + 层归一化

    这些就是下一篇的事了。

    参考链接

    • Bahdanau Attention 论文
    • Luong Attention 论文
    • Attention Is All You Need
    • PyTorch nn.MultiheadAttention 官方文档
    赞(0)
    未经允许不得转载:171主机测评 » 让模型学会“回头看”:Seq2Seq 与 Attention 机制的通俗图解与代码实战【NLP系列第三篇】
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址