| """
|
| Transformer Encoder 模块 — Person B 负责实现
|
|
|
| 包含:
|
| - TransformerEncoderLayer: 单层编码器
|
| - TransformerEncoder: 多层编码器堆叠
|
|
|
| 架构 (Pre-LayerNorm):
|
| x → LayerNorm → MultiHeadAttention → Residual → LayerNorm → FFN → Residual
|
| """
|
|
|
| from __future__ import annotations
|
|
|
| import copy
|
|
|
| import torch
|
| import torch.nn as nn
|
| from typing import Optional
|
|
|
| from easytranslate.model.attention import MultiHeadAttention, FlashMultiHeadAttention
|
|
|
|
|
| class TransformerEncoderLayer(nn.Module):
|
| """
|
| 单层 Transformer Encoder。
|
|
|
| 架构 (Pre-LayerNorm):
|
| x → LayerNorm → MultiHeadAttention → Residual → LayerNorm → FFN → Residual
|
| """
|
|
|
| def __init__(
|
| self,
|
| d_model: int = 512,
|
| nhead: int = 8,
|
| dim_feedforward: int = 2048,
|
| dropout: float = 0.1,
|
| activation: str = "gelu",
|
| use_flash_attention: bool = True,
|
| use_rotary_embedding: bool = True,
|
| pre_norm: bool = True,
|
| ):
|
| super().__init__()
|
| self.pre_norm = pre_norm
|
|
|
|
|
| attn_cls = FlashMultiHeadAttention if use_flash_attention else MultiHeadAttention
|
| self.self_attn = attn_cls(
|
| d_model=d_model,
|
| nhead=nhead,
|
| dropout=dropout,
|
| use_rotary_embedding=use_rotary_embedding,
|
| )
|
|
|
|
|
| self.linear1 = nn.Linear(d_model, dim_feedforward)
|
| self.activation = nn.GELU() if activation == "gelu" else nn.ReLU()
|
| self.dropout = nn.Dropout(p=dropout)
|
| self.linear2 = nn.Linear(dim_feedforward, d_model)
|
|
|
|
|
| self.norm1 = nn.LayerNorm(d_model)
|
| self.norm2 = nn.LayerNorm(d_model)
|
|
|
|
|
| self.dropout1 = nn.Dropout(p=dropout)
|
| self.dropout2 = nn.Dropout(p=dropout)
|
|
|
| def forward(
|
| self,
|
| src: torch.Tensor,
|
| src_key_padding_mask: Optional[torch.BoolTensor] = None,
|
| ) -> torch.Tensor:
|
| """Pre-LayerNorm 前向传播。"""
|
| if self.pre_norm:
|
|
|
| residual = src
|
| src = self.norm1(src)
|
| src = self.self_attn(
|
| src, src, src,
|
| key_padding_mask=src_key_padding_mask,
|
| )
|
| src = residual + self.dropout1(src)
|
|
|
|
|
| residual = src
|
| src = self.norm2(src)
|
| src = self.linear2(self.dropout(self.activation(self.linear1(src))))
|
| src = residual + self.dropout2(src)
|
| else:
|
|
|
| residual = src
|
| src = self.self_attn(
|
| src, src, src,
|
| key_padding_mask=src_key_padding_mask,
|
| )
|
| src = self.norm1(residual + self.dropout1(src))
|
|
|
| residual = src
|
| src = self.linear2(self.dropout(self.activation(self.linear1(src))))
|
| src = self.norm2(residual + self.dropout2(src))
|
|
|
| return src
|
|
|
|
|
| class TransformerEncoder(nn.Module):
|
| """
|
| 多层 Transformer Encoder。
|
| """
|
|
|
| def __init__(self, encoder_layer: TransformerEncoderLayer, num_layers: int):
|
| super().__init__()
|
| self.layers = nn.ModuleList(
|
| [copy.deepcopy(encoder_layer) for _ in range(num_layers)]
|
| )
|
| self.num_layers = num_layers
|
| self.norm = nn.LayerNorm(encoder_layer.self_attn.d_model)
|
|
|
| def forward(
|
| self,
|
| src: torch.Tensor,
|
| src_key_padding_mask: Optional[torch.BoolTensor] = None,
|
| ) -> torch.Tensor:
|
| output = src
|
| for layer in self.layers:
|
| output = layer(output, src_key_padding_mask=src_key_padding_mask)
|
| output = self.norm(output)
|
| return output
|
|
|