sdfjliom's picture
model (#4)
4d62693
Raw
History Blame Contribute Delete
3.85 kB
"""
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
# Self-Attention
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,
)
# Feed-Forward Network
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)
# LayerNorm
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
# Dropout for residual
self.dropout1 = nn.Dropout(p=dropout)
self.dropout2 = nn.Dropout(p=dropout)
def forward(
self,
src: torch.Tensor, # [B, S, D]
src_key_padding_mask: Optional[torch.BoolTensor] = None, # [B, S]
) -> torch.Tensor:
"""Pre-LayerNorm 前向传播。"""
if self.pre_norm:
# 1. Self-Attention sublayer
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)
# 2. FFN sublayer
residual = src
src = self.norm2(src)
src = self.linear2(self.dropout(self.activation(self.linear1(src))))
src = residual + self.dropout2(src)
else:
# Post-LayerNorm (备用)
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