| """ |
| 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 |
|
|