File size: 9,174 Bytes
c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 4d62693 c1a46f7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 | """
Transformer 翻译模型主体 — Person B 负责实现
这是整个模型的核心文件,定义 Encoder-Decoder 架构。
架构设计:
- Pre-LayerNorm Transformer (训练更稳定)
- 可选 Flash Attention 2 (加速注意力计算)
- 可选 RoPE 旋转位置编码 (替代传统正弦位置编码)
- 共享 Embedding 权重 (可选)
"""
from __future__ import annotations
import logging
import math
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from easytranslate.model.encoder import TransformerEncoder, TransformerEncoderLayer
from easytranslate.model.decoder import TransformerDecoder, TransformerDecoderLayer
from easytranslate.model.positional import SinusoidalPositionalEncoding, RotaryPositionalEmbedding
logger = logging.getLogger(__name__)
class TransformerTranslationModel(nn.Module):
"""
完整的 Transformer 英中翻译模型。
架构:
Source Embedding + Positional Encoding
→ Transformer Encoder (N layers)
→ Transformer Decoder (N layers)
→ Linear Projection → Softmax
TODO [Person B]: 实现以下内容:
__init__:
1. 源语言 Embedding: nn.Embedding(src_vocab_size, d_model)
2. 目标语言 Embedding: nn.Embedding(tgt_vocab_size, d_model)
3. 位置编码: SinusoidalPositionalEncoding 或 RotaryPositionalEmbedding
4. Transformer Encoder
5. Transformer Decoder
6. 输出投影层: nn.Linear(d_model, tgt_vocab_size)
7. (可选) 共享 target embedding 和 output projection 的权重
8. 调用 _init_weights() 初始化参数
forward:
1. 源语言 embedding + 位置编码 → encoder_input
2. 目标语言 embedding + 位置编码 → decoder_input
3. 生成 masks (src_key_padding_mask, tgt_key_padding_mask, tgt_mask)
4. encoder_output = encoder(encoder_input, src_key_padding_mask)
5. decoder_output = decoder(decoder_input, encoder_output, masks...)
6. logits = output_projection(decoder_output)
7. 返回 logits [B, T, tgt_vocab_size]
"""
def __init__(
self,
src_vocab_size: int,
tgt_vocab_size: int,
d_model: int = 512,
nhead: int = 8,
num_encoder_layers: int = 6,
num_decoder_layers: int = 6,
dim_feedforward: int = 2048,
dropout: float = 0.1,
activation: str = "gelu",
max_seq_len: int = 512,
use_flash_attention: bool = True,
use_rotary_embedding: bool = True,
pre_norm: bool = True,
pad_id: int = 0,
share_embedding: bool = False,
):
super().__init__()
self.d_model = d_model
self.pad_id = pad_id
self.use_rotary_embedding = use_rotary_embedding
self.max_seq_len = max_seq_len
# Embeddings
self.src_embed = nn.Embedding(src_vocab_size, d_model)
self.tgt_embed = nn.Embedding(tgt_vocab_size, d_model)
self.embed_scale = math.sqrt(d_model)
# Positional encoding
if use_rotary_embedding:
# RoPE 在 attention 内部应用到 Q/K,不需要额外的位置编码层
self.pos_encoding: Optional[nn.Module] = None
rope = RotaryPositionalEmbedding(
dim=d_model // nhead,
max_seq_len=max_seq_len,
)
else:
self.pos_encoding = SinusoidalPositionalEncoding(
d_model=d_model,
max_seq_len=max_seq_len,
dropout=dropout,
)
rope = None
# Encoder
encoder_layer = TransformerEncoderLayer(
d_model=d_model,
nhead=nhead,
dim_feedforward=dim_feedforward,
dropout=dropout,
activation=activation,
use_flash_attention=use_flash_attention,
use_rotary_embedding=use_rotary_embedding,
pre_norm=pre_norm,
)
self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers)
# Decoder
decoder_layer = TransformerDecoderLayer(
d_model=d_model,
nhead=nhead,
dim_feedforward=dim_feedforward,
dropout=dropout,
activation=activation,
use_flash_attention=use_flash_attention,
use_rotary_embedding=use_rotary_embedding,
pre_norm=pre_norm,
)
self.decoder = TransformerDecoder(decoder_layer, num_decoder_layers)
# 将 RoPE 注入到 encoder/decoder 的 attention 模块中
if rope is not None:
for layer in self.encoder.layers:
layer.self_attn.rope = rope
for layer in self.decoder.layers:
layer.self_attn.rope = rope
# cross-attention 不使用 RoPE
layer.multihead_attn.rope = None
# Output projection
self.output_projection = nn.Linear(d_model, tgt_vocab_size)
# 可选: 共享目标语言 embedding 和输出投影权重
self.share_embedding = share_embedding
if share_embedding:
self.output_projection.weight = self.tgt_embed.weight
self._init_weights()
def _init_weights(self):
"""参数初始化。"""
for p in self.parameters():
if p.dim() > 1:
nn.init.xavier_uniform_(p)
for module in self.modules():
if isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0, std=self.d_model ** -0.5)
elif isinstance(module, nn.LayerNorm):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def _generate_square_subsequent_mask(self, sz: int, device: torch.device) -> torch.Tensor:
"""生成因果注意力掩码 (causal mask)。"""
mask = torch.triu(torch.ones(sz, sz, device=device), diagonal=1)
mask = mask.masked_fill(mask == 1, float("-inf"))
return mask
def forward(
self,
src_ids: torch.Tensor, # [B, S]
tgt_input_ids: torch.Tensor, # [B, T]
src_padding_mask: Optional[torch.BoolTensor] = None, # [B, S]
tgt_padding_mask: Optional[torch.BoolTensor] = None, # [B, T]
) -> torch.Tensor:
"""
前向传播。
Returns:
logits: [B, T, tgt_vocab_size]
"""
# 1. Embedding
src_emb = self.src_embed(src_ids) * self.embed_scale # [B, S, D]
tgt_emb = self.tgt_embed(tgt_input_ids) * self.embed_scale # [B, T, D]
# 2. Positional encoding (如果不使用 RoPE)
if self.pos_encoding is not None:
src_emb = self.pos_encoding(src_emb)
tgt_emb = self.pos_encoding(tgt_emb)
# 3. Masks
if src_padding_mask is None:
src_padding_mask = src_ids.eq(self.pad_id)
if tgt_padding_mask is None:
tgt_padding_mask = tgt_input_ids.eq(self.pad_id)
tgt_seq_len = tgt_input_ids.size(1)
tgt_mask = self._generate_square_subsequent_mask(tgt_seq_len, tgt_input_ids.device)
# 4. Encoder
encoder_output = self.encoder(src_emb, src_key_padding_mask=src_padding_mask)
# 5. Decoder
decoder_output = self.decoder(
tgt_emb,
encoder_output,
tgt_mask=tgt_mask,
memory_key_padding_mask=src_padding_mask,
tgt_key_padding_mask=tgt_padding_mask,
)
# 6. Output projection
logits = self.output_projection(decoder_output)
return logits
@torch.no_grad()
def encode(self, src_ids: torch.Tensor, src_padding_mask: Optional[torch.BoolTensor] = None) -> torch.Tensor:
"""仅编码(用于推理时复用 encoder 输出)。"""
src_emb = self.src_embed(src_ids) * self.embed_scale
if self.pos_encoding is not None:
src_emb = self.pos_encoding(src_emb)
if src_padding_mask is None:
src_padding_mask = src_ids.eq(self.pad_id)
encoder_output = self.encoder(src_emb, src_key_padding_mask=src_padding_mask)
return encoder_output
@torch.no_grad()
def decode_step(
self,
tgt_input_ids: torch.Tensor,
encoder_output: torch.Tensor,
src_padding_mask: Optional[torch.BoolTensor] = None,
) -> torch.Tensor:
"""解码一步(用于自回归推理)。"""
tgt_emb = self.tgt_embed(tgt_input_ids) * self.embed_scale
if self.pos_encoding is not None:
tgt_emb = self.pos_encoding(tgt_emb)
tgt_seq_len = tgt_input_ids.size(1)
tgt_mask = self._generate_square_subsequent_mask(tgt_seq_len, tgt_input_ids.device)
decoder_output = self.decoder(
tgt_emb,
encoder_output,
tgt_mask=tgt_mask,
memory_key_padding_mask=src_padding_mask,
)
# 取最后一个 token 的 logits
logits = self.output_projection(decoder_output[:, -1, :])
return logits
def count_parameters(self) -> int:
"""返回可训练参数数量。"""
return sum(p.numel() for p in self.parameters() if p.requires_grad)
|