| """
|
| 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
|
|
|
|
|
| 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)
|
|
|
|
|
| if use_rotary_embedding:
|
|
|
| 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_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_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)
|
|
|
|
|
| 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
|
|
|
| layer.multihead_attn.rope = None
|
|
|
|
|
| self.output_projection = nn.Linear(d_model, tgt_vocab_size)
|
|
|
|
|
| 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,
|
| tgt_input_ids: torch.Tensor,
|
| src_padding_mask: Optional[torch.BoolTensor] = None,
|
| tgt_padding_mask: Optional[torch.BoolTensor] = None,
|
| ) -> torch.Tensor:
|
| """
|
| 前向传播。
|
|
|
| Returns:
|
| logits: [B, T, tgt_vocab_size]
|
| """
|
|
|
| src_emb = self.src_embed(src_ids) * self.embed_scale
|
| tgt_emb = self.tgt_embed(tgt_input_ids) * self.embed_scale
|
|
|
|
|
| if self.pos_encoding is not None:
|
| src_emb = self.pos_encoding(src_emb)
|
| tgt_emb = self.pos_encoding(tgt_emb)
|
|
|
|
|
| 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)
|
|
|
|
|
| encoder_output = self.encoder(src_emb, src_key_padding_mask=src_padding_mask)
|
|
|
|
|
| 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,
|
| )
|
|
|
|
|
| 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,
|
| )
|
|
|
|
|
| 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)
|
|
|