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