sdfjliom's picture
model (#4)
4d62693
Raw
History Blame Contribute Delete
9.17 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, 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)