lijn14
创建工程
c1a46f7
Raw
History Blame
5.16 kB
"""
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__()
# TODO [Person B]: 实现模型初始化
# 保存超参数
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, # [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]
"""
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)