当前位置: 首页 > news >正文

Transformer架构实战:从零理解自注意力机制到BERT/GPT实现差异

Transformer架构实战:从零理解自注意力机制到BERT/GPT实现差异

当你在Colab中运行第一个Transformer模型时,是否注意到BERT和GPT对同一段文本的处理方式截然不同?这种差异源于Transformer架构中编码器与解码器的精妙设计。本文将用PyTorch代码拆解自注意力机制的核心实现,并通过可视化工具揭示BERT和GPT在位置编码、注意力掩码等关键组件上的技术分叉。

1. 自注意力机制的数学本质与代码实现

自注意力机制的核心在于计算查询(Q)、键(K)、值(V)三个矩阵的交互。假设输入序列长度为n,嵌入维度为d,则计算过程可分解为:

import torch import torch.nn.functional as F def self_attention(Q, K, V, mask=None): # Q,K,V shape: (batch_size, seq_len, d_model) d_k = Q.size(-1) scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k)) if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) weights = F.softmax(scores, dim=-1) return torch.matmul(weights, V), weights

关键参数对比

参数BERT-base典型值GPT-3典型值
注意力头数1296
隐藏层维度76812288
层数1296

注意:实际实现中会采用多头注意力机制,每个头的维度通常为d_model // num_heads

在可视化注意力权重时,BERT的注意力模式通常呈现对角线分布(关注局部上下文),而GPT由于因果掩码的限制,只能关注当前位置之前的token。使用matplotlib可以直观展示这种差异:

def plot_attention(weights, title): import matplotlib.pyplot as plt plt.imshow(weights[0, 0].detach().numpy(), cmap='viridis') plt.colorbar() plt.title(title) plt.show()

2. 位置编码:让模型理解顺序的两种范式

Transformer架构抛弃RNN的循环结构后,必须显式注入位置信息。BERT和GPT采用了完全不同的策略:

  • BERT的绝对位置编码

    class BERTPositionalEncoding(nn.Module): def __init__(self, d_model, max_len=512): super().__init__() position = torch.arange(max_len).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)) pe = torch.zeros(max_len, d_model) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x): return x + self.pe[:x.size(1)]
  • GPT的相对位置编码: 现代GPT变体(如GPT-3)更多使用旋转位置编码(RoPE),其核心是:

    def apply_rotary_pos_emb(q, k, sin, cos): q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed

位置编码效果对比

特性绝对位置编码相对位置编码
最大长度限制有(如512)理论上无限
泛化能力对训练长度外位置泛化较差能更好处理长文本
计算复杂度O(1)O(n)

3. 架构差异:编码器与解码器的关键设计

BERT的编码器架构允许同时处理整个输入序列,而GPT的解码器必须遵循自回归生成模式。这种根本差异体现在:

BERT的编码器层实现

class BERTLayer(nn.Module): def __init__(self, d_model, num_heads): super().__init__() self.attention = MultiHeadAttention(d_model, num_heads) self.norm1 = nn.LayerNorm(d_model) self.ffn = PositionwiseFFN(d_model) self.norm2 = nn.LayerNorm(d_model) def forward(self, x, mask): attn_out, _ = self.attention(x, x, x, mask) x = self.norm1(x + attn_out) ffn_out = self.ffn(x) return self.norm2(x + ffn_out)

GPT的解码器层实现

class GPTLayer(nn.Module): def __init__(self, d_model, num_heads): super().__init__() self.masked_attention = MultiHeadAttention(d_model, num_heads) self.norm1 = nn.LayerNorm(d_model) self.cross_attention = MultiHeadAttention(d_model, num_heads) # 仅在有encoder时使用 self.norm2 = nn.LayerNorm(d_model) self.ffn = PositionwiseFFN(d_model) self.norm3 = nn.LayerNorm(d_model) def forward(self, x, causal_mask): attn_out, _ = self.masked_attention(x, x, x, causal_mask) x = self.norm1(x + attn_out) # 如果是encoder-decoder架构,此处会添加cross attention ffn_out = self.ffn(x) return self.norm3(x + ffn_out)

关键架构对比

  1. 注意力掩码

    • BERT使用padding mask处理变长输入
    • GPT必须额外使用因果掩码(causal mask)
    def create_causal_mask(size): mask = torch.triu(torch.ones(size, size), diagonal=1) return mask.masked_fill(mask == 1, float('-inf'))
  2. 训练目标

    • BERT采用MLM(掩码语言模型)目标
    • GPT采用标准语言模型目标(预测下一个token)

4. 实战对比:同一任务下的不同表现

在文本分类任务中,我们可以清晰观察到两种架构的差异。以IMDb影评分类为例:

BERT实现方案

from transformers import BertModel, BertTokenizer tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') model = BertModel.from_pretrained('bert-base-uncased') inputs = tokenizer("This movie was amazing!", return_tensors="pt") outputs = model(**inputs) cls_embedding = outputs.last_hidden_state[:, 0, :] # 取[CLS]标记对应的嵌入

GPT实现方案

from transformers import GPT2Model, GPT2Tokenizer tokenizer = GPT2Tokenizer.from_pretrained('gpt2') model = GPT2Model.from_pretrained('gpt2') inputs = tokenizer("This movie was amazing!", return_tensors="pt") outputs = model(**inputs) # 需要取最后一个token的嵌入或做pooling last_embedding = outputs.last_hidden_state[:, -1, :]

性能对比(IMDb测试集):

指标BERT-baseGPT-2
准确率92.3%88.7%
训练速度1.2小时0.8小时
显存占用3.2GB2.7GB

提示:实际应用中,GPT类模型需要通过添加分类头或prompt tuning来适配分类任务

在Colab实操中,可以通过以下代码可视化两者的注意力模式差异:

def compare_attention(text): bert_outs = bert_model(**bert_tokenizer(text, return_tensors='pt')).attentions gpt_outs = gpt_model(**gpt_tokenizer(text, return_tensors='pt')).attentions plot_attention(bert_outs[0][0], "BERT Attention") plot_attention(gpt_outs[0][0], "GPT Attention")
http://www.cnnetsun.cn/news/1365923.html

相关文章:

  • 3个关键问题:如何用TT-NN动态量化实现高效混合精度推理
  • Oni编辑器宏录制与回放:终极自动化重复任务指南
  • WAN2.2文生视频多场景落地:律所法律条款→情景剧式普法短视频自动生成
  • LSTM与RMBG-2.0结合的图像序列处理方案
  • Swin2SR快速部署指南:3步搭建个人图片修复工具
  • 终极指南:如何为你的Gridea静态博客构建坚不可摧的安全防线
  • RexUniNLU效果惊艳展示:中文长文本中嵌套事件与多跳关系精准抽取
  • 如何实现Giscus评论系统的实时数据同步:GitHub讨论更新与前端刷新机制详解
  • S12SD紫外线传感器原理与ESP32嵌入式集成指南
  • 7个关键指标!Walrus存储节点监控完整指南:确保去中心化存储高可用性
  • GPTs项目维护指南:5个关键策略确保长期可持续发展
  • TGN超参数调优终极指南:提升模型性能的10个关键技巧
  • 如何实现无障碍支持?Semi Design ARIA属性技术解析
  • Maccy更新失败解决指南:3种手动升级方法详解
  • 终极指南:如何使用pypdf提升PDF文档在Google中的排名
  • 如何在Robo 3T中配置MongoDB Atlas文本搜索索引:完整指南
  • 掌握ipatool日志系统:高效调试与问题追踪的完整指南
  • 终极窗口置顶解决方案:这款开源工具让你的工作窗口永不“失踪”
  • 什么是大模型?一文彻底搞懂大模型定义!!!未来淘汰你的不是AI,而是掌握了AI的人
  • OFA-large镜像保姆级部署教程:开箱即用跑通SNLI-VE语义蕴含任务
  • Qwen3-ASR-0.6B本地化部署实操:NVIDIA Jetson Orin边缘设备适配指南
  • Qwen3-TTS-Tokenizer-12Hz智能助手:嵌入式语音交互的轻量编码方案
  • 幻镜NEURAL MASK在IP形象开发中的应用:角色素材标准化生产流程
  • MGeo地址解析开源模型部署实操:Ubuntu/CentOS环境Gradio服务一键启动
  • Qwen3.5-35B-A3B-AWQ-4bit镜像快速上手:无需conda/pip,直接supervisorctl启动
  • 嵌入式菜鸟的进阶之路——C的输入输出
  • SOONet视频预处理指南:FFmpeg抽帧/重编码/分辨率适配最佳实践
  • sse哈工大C语言编程练习47
  • Cosmos-Reason1-7B应用场景:智能制造产线视频中设备异常振动与故障关联分析
  • 南北阁 Nanbeige 4.1-3B 效果惊艳:中文法律条文解读+案例匹配真实输出