Seq2Seq - 编码器(Encoder)和解码器(Decoder)

本节实现一个简单的 Seq2Seq(Sequence to Sequence)模型 的编码器(Encoder)和解码器(Decoder)部分。

重点把握Seq2Seq 模型的整体工作流程

理解编码器(Encoder)和解码器(Decoder)代码

本小节引入了nn.GRU API的调用,nn.GRU具体参数将在下一小节进行补充讲解

1. 编码器(Encoder

类定义
复制代码
class Encoder(nn.Module):
    def __init__(self, vocab_size, embedding_dim, hidden_size):
        super().__init__()
        self.emb = nn.Embedding(vocab_size, embedding_dim)
        self.rnn = nn.GRU(embedding_dim, hidden_size, batch_first=True)
  • vocab_size:输入词汇表的大小,即输入序列中可能出现的不同单词或标记的数量。

  • embedding_dim:嵌入层的维度,即每个单词或标记被映射到的向量空间的维度。

  • hidden_size:GRU(门控循环单元)的隐藏状态维度,决定了模型的内部状态大小。

主要组件
  1. 嵌入层(nn.Embedding

    • 嵌入层会将输入序列形状转换为 [batch_size, seq_len, embedding_dim] 的张量。

    • 这种映射是通过学习嵌入矩阵实现的,每个单词索引对应嵌入矩阵中的一行。

  2. GRU(nn.GRU

    • embedding_dim 是 GRU 的输入维度,hidden_size 是隐藏状态的维度。

    • batch_first=True 表示输入和输出的张量的第一个维度是批量大小(batch_size),而不是序列长度(seq_len)。

前向传播(forward
复制代码
def forward(self, x):
    embs = self.emb(x) #batch * token * embedding_dim
    gru_out, hidden = self.rnn(embs) #batch * token * hidden_size

    return gru_out, hidden
  • 输入 x 是一个形状为 [batch_size, seq_len] 的张量,表示一个批次的输入序列。

  • embs 是嵌入层的输出,形状为 [batch_size, seq_len, embedding_dim]

  • gru_out 是 GRU 的输出,形状为 [batch_size, seq_len, hidden_size],表示每个时间步的隐藏状态。

  • hidden 是 GRU 的最终隐藏状态,形状为 [1, batch_size, hidden_size],用于传递给解码器。

2. 解码器(Decoder)

类定义
复制代码
class Decoder(nn.Module):
    def __init__(self, vocab_size, embedding_dim, hidden_size):
        super().__init__()
        self.emb = nn.Embedding(vocab_size, embedding_dim)
        self.rnn = nn.GRU(embedding_dim, hidden_size, batch_first=True)
  • 解码器的结构与编码器类似,但它的作用是将编码器生成的上下文向量(hidden)解码为目标序列。
主要组件
  1. 嵌入层(nn.Embedding

    • 与编码器类似,将目标序列的单词索引映射到嵌入向量。
  2. GRU(nn.GRU

    • 与编码器中的 GRU 类似,但其输入是目标序列的嵌入向量,初始隐藏状态是编码器的最终隐藏状态。
前向传播(forward
复制代码
def forward(self, x, hx):
    embs = self.emb(x)
    gru_out, hidden = self.rnn(embs, hx=hx) #batch * token * hidden_size
    # batch * token * hidden_size
    # 1 * token * hidden_size

    return gru_out, hidden
  • 输入 x 是目标序列的单词索引,形状为 [batch_size, seq_len]

  • hx 是编码器的最终隐藏状态,形状为 [1, batch_size, hidden_size],作为解码器的初始隐藏状态。

  • embs 是目标序列的嵌入向量,形状为 [batch_size, seq_len, embedding_dim]

  • gru_out 是解码器 GRU 的输出,形状为 [batch_size, seq_len, hidden_size]

  • hidden 是解码器 GRU 的最终隐藏状态,形状为 [1, batch_size, hidden_size]

3. Seq2Seq 模型的整体工作流程⭐

  1. 编码阶段

    • 输入序列通过编码器的嵌入层,将单词索引映射为嵌入向量。

    • 嵌入向量通过 GRU,生成每个时间步的隐藏状态和最终的隐藏状态(上下文向量)。

    • 最终隐藏状态(hidden)作为编码器的输出,传递给解码器。

  2. 解码阶段

    • 解码器的初始隐藏状态是编码器的最终隐藏状态。

    • 解码器逐个生成目标序列的单词,每次生成一个单词后,将该单词的嵌入向量作为下一次输入,同时更新隐藏状态。

    • 通过这种方式,解码器逐步生成目标序列。

相关推荐
小刘学技术几秒前
AI人工智能中的类别不平衡问题:成因、影响与解决方案
开发语言·人工智能·python·机器学习
阿拉雷️几秒前
部署实战】Docker + AI Agent:让AI一键部署Spring Boot到服务器,从打包到上线只要一条指令
人工智能·spring boot·docker
饼干哥哥4 分钟前
我用千问3.8跑通了Reddit自动海外获客部门,成本砍 10 倍!
人工智能·开源·创业
武子康9 分钟前
Claude Code 权限分析器为什么必须 Fail Closed:v2.1.214 暴露的 5 类边界 + 6 类不能推出的结论
人工智能·ai编程·claude
元直数字电路验证10 分钟前
深入理解 AI Agent:从模型能力到生产级系统的完整路线图
人工智能·langchain·aigc·agent·智能体
zyplayer-doc15 分钟前
研发接口文档怎么长期维护:zyplayer-doc把API、Markdown和变更记录放进同一个知识库
大数据·数据库·人工智能·笔记·pdf·ocr
赋创小助手23 分钟前
AMD Helios AI机架技术详解:72颗MI455X、EPYC Venice与UALoE架构
人工智能·架构·amd·amd helios ai机架·mi455x·epyc venice·ualoe架构
动物园猫25 分钟前
PCB表面缺陷目标检测数据集:6类别、3,500张图像 | 目标检测
人工智能·目标检测·计算机视觉
也非非也25 分钟前
Agent支付的真正战争,不在演示台,而在后台
人工智能·ai编程·vibecoding·waic
义嘉泰29 分钟前
国产 eMMC 替代选型:XTX XT28EG08GA5SL / XT28EG16GA5SL 解析
人工智能·科技·芯片