尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Transformer架构与自注意力机制实现详解

Transformer架构与自注意力机制实现详解 1. Transformer架构与注意力机制深度解析在深度学习领域Transformer模型彻底改变了序列建模的范式。与传统的RNN和CNN不同Transformer通过自注意力机制实现了对序列数据的并行处理显著提升了模型效率和性能表现。1.1 自注意力机制核心原理自注意力机制的核心在于建立序列元素间的动态关联。给定输入序列X∈ℝ^(n×d)模型通过三个可学习的权重矩阵WQ、WK、WV∈ℝ^(d×d)分别生成查询向量Q XWQ键向量K XWK值向量V XWV注意力得分的计算采用缩放点积形式def scaled_dot_product_attention(Q, K, V, maskNone): d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) p_attn torch.softmax(scores, dim-1) return torch.matmul(p_attn, V), p_attn这种设计使得模型能够动态关注不同位置的元素解决了传统RNN的长距离依赖问题。1.2 多头注意力实现细节多头注意力将输入分割到多个子空间并行计算class MultiHeadAttention(nn.Module): def __init__(self, h, d_model, dropout0.1): super().__init__() assert d_model % h 0 self.d_k d_model // h self.h h self.linears nn.ModuleList([nn.Linear(d_model, d_model) for _ in range(4)]) self.dropout nn.Dropout(pdropout) def forward(self, query, key, value, maskNone): nbatches query.size(0) # 线性变换并分割头 query, key, value [ lin(x).view(nbatches, -1, self.h, self.d_k).transpose(1, 2) for lin, x in zip(self.linears, (query, key, value)) ] # 计算注意力 x, self.attn scaled_dot_product_attention( query, key, value, maskmask, dropoutself.dropout ) # 合并头并做最终线性变换 x x.transpose(1, 2).contiguous().view(nbatches, -1, self.h * self.d_k) return self.linears[-1](x)每个头学习不同的注意力模式最后将各头的输出拼接并通过线性层融合增强了模型的表达能力。2. Transformer核心组件实现2.1 编码器层设计编码器层包含两个核心子层class EncoderLayer(nn.Module): def __init__(self, size, self_attn, feed_forward, dropout): super().__init__() self.self_attn self_attn self.feed_forward feed_forward self.sublayer nn.ModuleList([ SublayerConnection(size, dropout) for _ in range(2) ]) self.size size def forward(self, x, mask): x self.sublayer[0](x, lambda x: self.self_attn(x, x, x, mask)) return self.sublayer[1](x, self.feed_forward)其中子层连接实现了残差连接和层归一化class SublayerConnection(nn.Module): def __init__(self, size, dropout): super().__init__() self.norm LayerNorm(size) self.dropout nn.Dropout(dropout) def forward(self, x, sublayer): return x self.dropout(sublayer(self.norm(x)))2.2 解码器层特殊设计解码器层在编码器基础上增加了交叉注意力class DecoderLayer(nn.Module): def __init__(self, size, self_attn, src_attn, feed_forward, dropout): super().__init__() self.size size self.self_attn self_attn self.src_attn src_attn self.feed_forward feed_forward self.sublayer nn.ModuleList([ SublayerConnection(size, dropout) for _ in range(3) ]) def forward(self, x, memory, src_mask, tgt_mask): x self.sublayer[0](x, lambda x: self.self_attn(x, x, x, tgt_mask)) x self.sublayer[1](x, lambda x: self.src_attn(x, memory, memory, src_mask)) return self.sublayer[2](x, self.feed_forward)掩码自注意力确保解码时只能看到当前位置之前的标记这是实现自回归生成的关键。3. 位置编码与词嵌入3.1 位置编码数学原理位置编码使用不同频率的正余弦函数class PositionalEncoding(nn.Module): def __init__(self, d_model, dropout, max_len5000): super().__init__() self.dropout nn.Dropout(pdropout) pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1) div_term torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000.0) / d_model)) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) pe pe.unsqueeze(0) self.register_buffer(pe, pe) def forward(self, x): x x self.pe[:, :x.size(1)] return self.dropout(x)这种设计使模型能够学习到相对位置关系且对任意长度的序列都有良好的泛化能力。3.2 词嵌入实现技巧词嵌入层将离散的token转换为连续向量class Embeddings(nn.Module): def __init__(self, d_model, vocab): super().__init__() self.lut nn.Embedding(vocab, d_model) self.d_model d_model def forward(self, x): return self.lut(x) * math.sqrt(self.d_model)乘以√d_model是为了保持嵌入值与位置编码相加后的数值稳定性。4. 完整Transformer组装4.1 模型构建流程def make_model(src_vocab, tgt_vocab, N6, d_model512, d_ff2048, h8, dropout0.1): c copy.deepcopy attn MultiHeadedAttention(h, d_model) ff PositionwiseFeedForward(d_model, d_ff, dropout) position PositionalEncoding(d_model, dropout) model Transformer( Encoder(EncoderLayer(d_model, c(attn), c(ff), dropout), N), Decoder(DecoderLayer(d_model, c(attn), c(attn), c(ff), dropout), N), nn.Sequential(Embeddings(d_model, src_vocab), c(position)), nn.Sequential(Embeddings(d_model, tgt_vocab), c(position)), Generator(d_model, tgt_vocab)) # 参数初始化 for p in model.parameters(): if p.dim() 1: nn.init.xavier_uniform_(p) return model关键参数说明N编码器/解码器层数默认6d_model模型维度默认512d_ff前馈网络隐藏层维度默认2048h注意力头数默认84.2 训练技巧与参数设置实际训练时需要注意学习率调度使用warmup策略class NoamOpt: def __init__(self, model_size, factor, warmup, optimizer): self.optimizer optimizer self._step 0 self.warmup warmup self.factor factor self.model_size model_size self._rate 0 def step(self): self._step 1 rate self.rate() for p in self.optimizer.param_groups: p[lr] rate self._rate rate self.optimizer.step() def rate(self, stepNone): if step is None: step self._step return self.factor * \ (self.model_size ** (-0.5) * min(step ** (-0.5), step * self.warmup ** (-1.5)))标签平滑提升模型泛化能力class LabelSmoothing(nn.Module): def __init__(self, size, padding_idx, smoothing0.0): super().__init__() self.criterion nn.KLDivLoss(reductionsum) self.padding_idx padding_idx self.confidence 1.0 - smoothing self.smoothing smoothing self.size size self.true_dist None def forward(self, x, target): assert x.size(1) self.size true_dist x.data.clone() true_dist.fill_(self.smoothing / (self.size - 2)) true_dist.scatter_(1, target.data.unsqueeze(1), self.confidence) true_dist[:, self.padding_idx] 0 mask torch.nonzero(target.data self.padding_idx) if mask.dim() 0: true_dist.index_fill_(0, mask.squeeze(), 0.0) self.true_dist true_dist return self.criterion(x, true_dist)5. 实战英法翻译系统构建5.1 数据处理流程使用子词分词from transformers import XLMTokenizer tokenizer XLMTokenizer.from_pretrained(xlm-clm-enfr-1024) en_text I dont speak French. fr_text Je ne parle pas français. en_tokens tokenizer.tokenize(en_text) # [i/w, don/w, t/w, ...] fr_tokens tokenizer.tokenize(fr_text) # [je/w, ne/w, parle/w, ...]构建词汇表from collections import Counter def build_vocab(token_lists, max_size50000): counter Counter() for tokens in token_lists: counter.update(tokens) vocab {word:i2 for i, (word,_) in enumerate(counter.most_common(max_size))} vocab[pad] 0 vocab[unk] 1 return vocab en_vocab build_vocab(en_tokenized) fr_vocab build_vocab(fr_tokenized)5.2 模型训练关键步骤# 初始化模型 model make_model(len(en_vocab), len(fr_vocab)) model.to(device) # 定义优化器和损失函数 optimizer NoamOpt(model.src_embed[0].d_model, 2, 4000, torch.optim.Adam(model.parameters(), lr0, betas(0.9, 0.98), eps1e-9)) criterion LabelSmoothing(sizelen(fr_vocab), padding_idx0, smoothing0.1) # 训练循环 for epoch in range(epochs): model.train() for batch in train_loader: src batch.en.to(device) trg batch.fr.to(device) # 前向传播 out model(src, trg[:, :-1]) loss criterion(out.contiguous().view(-1, out.size(-1)), trg[:, 1:].contiguous().view(-1)) # 反向传播 optimizer.optimizer.zero_grad() loss.backward() optimizer.step()6. 性能优化技巧混合精度训练scaler torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): out model(src, trg[:, :-1]) loss criterion(out.contiguous().view(-1, out.size(-1)), trg[:, 1:].contiguous().view(-1)) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)批处理技巧动态padding同批次样本padding到相同长度桶排序将长度相近的样本放在同批次7. 模型评估与推理7.1 评估指标计算使用BLEU分数评估翻译质量from nltk.translate.bleu_score import corpus_bleu def evaluate(model, val_loader, fr_vocab): model.eval() refs [] hyps [] with torch.no_grad(): for batch in val_loader: src batch.en.to(device) trg batch.fr.to(device) # 生成翻译 preds greedy_decode(model, src, max_len100) # 转换为文本 ref_texts [[fr_idx_dict[idx] for idx in seq if idx not in (0,1,2)] for seq in trg.cpu().numpy()] hyp_texts [[fr_idx_dict[idx] for idx in seq if idx not in (0,1,2)] for seq in preds.cpu().numpy()] refs.extend([[ref] for ref in ref_texts]) hyps.extend(hyp_texts) return corpus_bleu(refs, hyps)7.2 贪心解码实现def greedy_decode(model, src, max_len, start_symbol2): memory model.encode(src, None) ys torch.ones(1, 1).fill_(start_symbol).type_as(src.data) for i in range(max_len-1): out model.decode(memory, None, ys, subsequent_mask(ys.size(1)).type_as(src.data)) prob model.generator(out[:, -1]) _, next_word torch.max(prob, dim1) next_word next_word.data[0] ys torch.cat([ys, torch.ones(1, 1).type_as(src.data).fill_(next_word)], dim1) if next_word 3: # EOS token break return ys8. 常见问题排查训练不收敛检查梯度流动各层梯度值应在合理范围验证损失计算确保padding部分被正确mask调整学习率使用warmup策略过拟合问题增加dropout率使用更激进的标签平滑添加更多训练数据推理结果异常检查解码温度设置验证词汇表映射是否正确确保输入序列长度不超过模型最大位置编码在实际项目中我发现模型对长序列的处理能力与位置编码设计密切相关。当输入序列超过训练时的最大长度时性能会显著下降。解决方案是预训练时使用足够大的max_len参数或者在微调阶段重新初始化位置编码。
返回列表