超详细!Transformer 模型解码器及核心模块实现全梳理

诗人啊_程序员 2025-08-16 01:34:23

一、解码器层(DecoderLayer)—— 单层级核心运算

核心代码

python

class DecoderLayer(nn.Module):
    def __init__(self, size, adn, source, feed_forward, dropout):
        super().__init__()
        self.size = size  # 自注意力维度
        self.adn = adn  # 多头自注意力(q=k=v)
        self.source = source  # 编码器-解码器注意力(k=v用memory)
        self.feed_forward = feed_forward  # 前馈网络
        self.sublayers = clones(SublayerConnection(size, dropout), 3)  # 残差+归一化封装

    def forward(self, x, memory, source_mask, target_mask):
        # 1. 自注意力(防未来信息泄露,用target_mask)
        x = self.sublayers[0](x, lambda x: self.adn(x, x, x, target_mask))
        # 2. 编码器-解码器注意力(对齐源文本,用source_mask)
        x = self.sublayers[1](x, lambda x: self.source(x, memory, memory, source_mask))
        # 3. 前馈网络(维度变换+非线性)
        return self.sublayers[2](x, self.feed_forward)

关键逻辑

  • 掩码作用target_mask 遮掉后续词(翻译时不能偷看未来内容),source_mask 忽略源文本里的 padding。
  • 执行顺序:必须先自注意力 “理解当前生成内容”,再用交叉注意力 “对齐源文本信息”,最后前馈网络强化特征。

二、解码器(Decoder)—— 多层堆叠迭代优化

核心代码

python

class Decoder(nn.Module):
    def __init__(self, layer, n):
        super().__init__()
        self.layers = clones(layer, n)  # 克隆n层(通常n=6)
        self.norm = LayerNorm(layer.size)  # 最终归一化

    def forward(self, x, memory, source_mask, target_mask):
        for layer in self.layers:
            x = layer(x, memory, source_mask, target_mask)  # 逐层加工
        return self.norm(x)  # 统一输出分布,方便后续计算

核心要点

  • 层数对齐:解码器层数要和编码器一致(都是 6 层),保证特征交互充分。
  • 输入变化:每层都用同一个 memory(编码器输出),但 x 会随着层堆叠不断优化,像 “草稿→精修” 的过程。

三、输出层(Generator)—— 转成词表概率

核心代码

python

class Generator(nn.Module):
    def __init__(self, d_model, vocab_size):
        super().__init__()
        self.linear = nn.Linear(d_model, vocab_size)  # 把解码器输出维度转成词表大小

    def forward(self, x):
        # log_softmax方便算交叉熵损失,输出每个位置的词概率
        return F.log_softmax(self.linear(x), dim=-1)

关键细节

  • 输入是解码器的输出(形状 [batch, seq_len, d_model]),输出是 [batch, seq_len, vocab_size],每个位置对应词表所有词的概率。
  • 用 log_softmax 而不是普通 softmax,数值计算更稳定,训练时算损失更方便。

四、模型串联(Encoder-Decoder)—— 从输入到输出的完整流程

核心代码

python

class EncoderDecoder(nn.Module):
    def __init__(self, encoder, decoder, src_embed, tgt_embed, generator):
        super().__init__()
        self.encoder = encoder  # 编码器(处理源文本)
        self.decoder = decoder  # 解码器(生成目标文本)
        self.src_embed = src_embed  # 源文本:词嵌入+位置编码
        self.tgt_embed = tgt_embed  # 目标文本:词嵌入+位置编码
        self.generator = generator  # 输出层转概率

    def forward(self, src, tgt, src_mask, tgt_mask):
        # 1. 编码源文本
        memory = self.encode(src, src_mask)
        # 2. 解码并生成结果
        return self.decode(memory, src_mask, tgt, tgt_mask)

    def encode(self, src, src_mask):
        return self.encoder(self.src_embed(src), src_mask)

    def decode(self, memory, src_mask, tgt, tgt_mask):
        return self.decoder(self.tgt_embed(tgt), memory, src_mask, tgt_mask)

执行流程

  1. 源文本处理src 先过 src_embed(词嵌入 + 位置编码),再丢给编码器生成 memory(理解后的源文本特征)。
  2. 目标文本生成tgt 过 tgt_embed 后,带着 memory 进解码器,逐层优化后,用 generator 转成词概率。

五、必懂工具函数

1. clones 函数(层克隆)

python

def clones(module, N):
    return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])

 

  • 作用:批量复制层(比如 6 个解码器层),每个层参数独立,避免互相干扰。

2. nn.Sequential vs nn.ModuleList

工具特点用在哪?
nn.Sequential层会按顺序执行,输入输出必须 “对上维度”固定流程(如 词嵌入→位置编码)
nn.ModuleList只存层,不自动执行需灵活调用(如解码器的多层迭代)

总结:Transformer 解码器核心逻辑

  • 流程:原始文本 → 嵌入 + 位置编码 → 编码器 → 解码器(多层自注意力 + 交叉注意力) → 输出层 → 词概率。
  • 关键:掩码防作弊(target_mask 遮未来,source_mask 跳 padding)、维度要对齐(d_model 贯穿全程)、多层迭代优化(解码器堆 6 层打磨结果)。

 

想自己实现 Transformer?照着这个结构,把编码器、注意力机制补上,就能跑通完整流程啦~ 有疑问评论区见!

...全文
101 回复 打赏 收藏 转发到动态 举报
写回复
用AI写文章
回复
切换为时间正序
请发表友善的回复…
发表回复
内容概要:SSD2828QN4是一款MIPI主桥接芯片,用于连接应用处理器与传统并行LCD接口及支持MIPI从属接口的LCD驱动器。该芯片支持最高每通道1Gbps的串行链路速度,最多可配置4个数据通道,显著减少了信号数量。它支持多种接口模式,包括RGB+SPI组合接口,适用于驱动智能或非智能显示面板,并能通过命令模式和视频模式传输数据。芯片内置时钟和复位模块、外部接口、协议控制单元(PCU)、包处理单元(PPU)、错误校正码/循环冗余校验(ECC/CRC)模块、长包和命令缓冲区、D-PHY控制器、模拟收发器以及内部锁相环(PLL),确保了高效的数据传输和系统稳定性。此外,文档详细描述了芯片的引脚分配、寄存器设置、操作模式、电源序列、时序特性等关键参数,为开发者提供了面的技术指导。 适合人群:具备一定硬件设计基础,从事嵌入式系统开发、显示技术研究的研发人员。 使用场景及目标:①实现应用处理器与MIPI兼容显示屏之间的高速数据传输;②优化显示系统的功耗表现,减少电磁干扰(EMI);③通过灵活配置不同接口模式来适应各种显示设备的需求。 阅读建议:此文档面向具有一定电子工程背景的专业人士,建议读者结合实际项目需求深入理解各章节内容,特别是关于寄存器配置、时序要求等方面的具体说明。对于初次接触此类技术的开发者而言,建议先熟悉基本概念再逐步掌握高级功能的应用方法。

12,046

社区成员

发帖
与我相关
我的任务
社区描述
创建由Python学习者和社区专家组成的国内最大的第三方Python中文社区,帮助社区成员更好地入门学习、职业成长和应用实践
python学习 企业社区
社区管理员
  • Python全栈技术社区
  • Lumos_zbj
  • 北侠大卫
加入社区
  • 近7日
  • 近30日
  • 至今
社区公告

创建由Python学习者和社区专家组成的国内最大的第三方Python中文社区,帮助社区成员更好地入门学习、职业成长和应用实践

  • 这里有最新最全的 Python 学习内容及资源,每月多达4次技术公开课
  • 这里有众多 Python 学习者,陪伴你一起交流成长
  • 这里有专业 Python 社区专家、讲师,帮助你跨越学习瓶颈,解决实操难题
  • 这里有丰富的社区活动,可以开阔眼界,结识更多同伴

【最新活动】:

  1. 周四技术公开课讲师招募中,点击查看详情
  2. “Python 社区专家团” 招募中,点击查看详情

 

试试用AI创作助手写篇文章吧