| """ |
| 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 |
| from easytranslate.model.decoder import TransformerDecoder |
| 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 |
|
|
| raise NotImplementedError("TODO: Person B 实现 TransformerTranslationModel.__init__") |
|
|
| def _init_weights(self): |
| """ |
| 参数初始化。 |
| |
| TODO [Person B]: 实现 Xavier/Kaiming 初始化: |
| - Embedding: normal_(0, d_model^-0.5) |
| - Linear: xavier_uniform_ |
| - LayerNorm: ones_ / zeros_ |
| """ |
| raise NotImplementedError("TODO: Person B 实现 _init_weights") |
|
|
| def _generate_square_subsequent_mask(self, sz: int, device: torch.device) -> torch.Tensor: |
| """ |
| 生成因果注意力掩码 (causal mask)。 |
| |
| TODO [Person B]: |
| 返回上三角矩阵 mask,shape [sz, sz], |
| mask[i][j] = -inf if j > i else 0 |
| """ |
| raise NotImplementedError("TODO: Person B 实现 _generate_square_subsequent_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] |
| """ |
| raise NotImplementedError("TODO: Person B 实现 forward") |
|
|
| @torch.no_grad() |
| def encode(self, src_ids: torch.Tensor, src_padding_mask: Optional[torch.BoolTensor] = None) -> torch.Tensor: |
| """ |
| 仅编码(用于推理时复用 encoder 输出)。 |
| |
| TODO [Person B]: |
| 1. src embedding + positional encoding |
| 2. encoder forward |
| 3. 返回 encoder_output |
| """ |
| raise NotImplementedError("TODO: Person B 实现 encode") |
|
|
| @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: |
| """ |
| 解码一步(用于自回归推理)。 |
| |
| TODO [Person B]: |
| 1. tgt embedding + positional encoding |
| 2. 生成 causal mask |
| 3. decoder forward |
| 4. 取最后一个 token 的 logits |
| 5. 返回 logits [B, vocab_size] |
| """ |
| raise NotImplementedError("TODO: Person B 实现 decode_step") |
|
|
| def count_parameters(self) -> int: |
| """返回可训练参数数量。""" |
| return sum(p.numel() for p in self.parameters() if p.requires_grad) |
|
|